Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
37 changes: 34 additions & 3 deletions eval/benchmarks/memops/scripts/run_memop_memory_baseline.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,10 @@

from eval.methods.shared.memqa_io import query_prompt_for_style # noqa: E402
from eval.methods.shared.memory_contract import QueryResult # noqa: E402
from metis.generation_utils import ( # noqa: E402
decode_generated_response,
resolve_generation_special_token_ids,
)
from eval.methods.shared.metis_loader import (
_chat_template as metis_chat_template,
infer_input_device,
Expand Down Expand Up @@ -152,15 +156,22 @@ def __init__(
low_rank_rank: int | None = None,
low_rank_policy: str = "after_each_commit",
low_rank_target: str = "state",
model_path: str = "",
):
self.checkpoint = Path(checkpoint).expanduser().resolve()
self.device = device
self.max_new_tokens = max_new_tokens
self.query_style = query_style
self.base_model_override = (
str(Path(model_path).expanduser().resolve())
if model_path
else None
)
self.model, self.tokenizer, self.load_report = load_v2_full_checkpoint(
self.checkpoint,
device=device,
dtype=metis_parse_dtype(dtype),
model_path=model_path or None,
device_map=device_map,
model_parallel_devices=model_parallel_devices,
max_memory=parse_max_memory(max_memory),
Expand Down Expand Up @@ -212,18 +223,22 @@ def query(self, question: str) -> QueryResult:
add_generation_prompt=True,
device=self.input_device,
)
eos_token_id, pad_token_id = resolve_generation_special_token_ids(
self.model,
self.tokenizer,
)
output_ids = self.model.generate(
input_ids=prefix_ids,
attention_mask=torch.ones_like(prefix_ids),
max_new_tokens=self.max_new_tokens,
do_sample=False,
use_cache=True,
pad_token_id=self.tokenizer.pad_token_id or self.tokenizer.eos_token_id,
eos_token_id=self.tokenizer.eos_token_id,
pad_token_id=pad_token_id,
eos_token_id=eos_token_id,
)
new_ids = output_ids[:, prefix_ids.shape[1] :]
return QueryResult(
raw_output=self.tokenizer.decode(new_ids[0], skip_special_tokens=True).strip(),
raw_output=decode_generated_response(self.tokenizer, new_ids[0]),
prompt_tokens=int(prefix_ids.shape[1]),
latency_sec=round(time.time() - started, 3),
debug={
Expand Down Expand Up @@ -293,6 +308,7 @@ def build_baseline(args: argparse.Namespace) -> tuple[Any, str, str]:
args.metis_low_rank_rank,
args.metis_low_rank_policy,
args.metis_low_rank_target,
model_path=args.model_path,
),
str(Path(args.checkpoint).expanduser().resolve()),
"metis_memory_state_only",
Expand Down Expand Up @@ -329,6 +345,11 @@ def output_record(
"baseline": args.method,
"model_label": args.model_label,
"model_path": model_path,
"base_model_override": (
getattr(args, "metis_base_model_override", None)
if args.method == "metis"
else None
),
"adapter_dir": args.adapter_dir if args.method == "delta_mem" else None,
"checkpoint": args.checkpoint if args.method == "metis" else None,
"device": args.device,
Expand Down Expand Up @@ -361,6 +382,11 @@ def run(args: argparse.Namespace) -> Path:
baseline, model_path, context_policy = build_baseline(args)
args.metis_input_device = str(getattr(baseline, "input_device", args.device))
args.metis_load_report = getattr(baseline, "load_report", None)
args.metis_base_model_override = getattr(
baseline,
"base_model_override",
None,
)
dataset_label = instances[0].get("dataset", "memop") if instances else "memop"
meta = {
"run_id": args.run_id,
Expand All @@ -371,6 +397,11 @@ def run(args: argparse.Namespace) -> Path:
"baseline": args.method,
"model_label": args.model_label,
"model_path": model_path,
"base_model_override": (
args.metis_base_model_override
if args.method == "metis"
else None
),
"adapter_dir": args.adapter_dir if args.method == "delta_mem" else None,
"checkpoint": args.checkpoint if args.method == "metis" else None,
"instances": str(args.instances),
Expand Down
14 changes: 11 additions & 3 deletions eval/benchmarks/memops/scripts/run_qwen_plain_context.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,11 @@
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer

from metis.generation_utils import (
decode_generated_response,
resolve_generation_special_token_ids,
)


def utc_now() -> str:
return dt.datetime.now(dt.timezone.utc).isoformat()
Expand Down Expand Up @@ -130,17 +135,20 @@ def generate(
for key, value in list(encoded.items()):
if torch.is_tensor(value) and value.ndim == 2 and value.shape[1] == original_prompt_tokens:
encoded[key] = value[:, -max_input_tokens:]
eos_token_id, pad_token_id = resolve_generation_special_token_ids(
model, tokenizer
)
output_ids = model.generate(
**encoded,
max_new_tokens=max_new_tokens,
do_sample=False,
use_cache=True,
pad_token_id=tokenizer.pad_token_id or tokenizer.eos_token_id,
eos_token_id=tokenizer.eos_token_id,
pad_token_id=pad_token_id,
eos_token_id=eos_token_id,
)
new_ids = output_ids[:, encoded["input_ids"].shape[1] :]
return (
tokenizer.decode(new_ids[0], skip_special_tokens=True).strip(),
decode_generated_response(tokenizer, new_ids[0]),
int(encoded["input_ids"].shape[1]),
round(time.time() - started, 3),
original_prompt_tokens,
Expand Down
14 changes: 11 additions & 3 deletions eval/benchmarks/memqa/scripts/run_base_context.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,11 @@
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer

from metis.generation_utils import (
decode_generated_response,
resolve_generation_special_token_ids,
)


def utc_now() -> str:
return dt.datetime.now(dt.timezone.utc).isoformat()
Expand Down Expand Up @@ -112,16 +117,19 @@ def generate(
started = time.time()
text = chat_text(tokenizer, prompt)
encoded = tokenizer(text, return_tensors="pt", add_special_tokens=False).to(device)
eos_token_id, pad_token_id = resolve_generation_special_token_ids(
model, tokenizer
)
output_ids = model.generate(
**encoded,
max_new_tokens=max_new_tokens,
do_sample=False,
use_cache=True,
pad_token_id=tokenizer.pad_token_id or tokenizer.eos_token_id,
eos_token_id=tokenizer.eos_token_id,
pad_token_id=pad_token_id,
eos_token_id=eos_token_id,
)
new_ids = output_ids[:, encoded.input_ids.shape[1] :]
return tokenizer.decode(new_ids[0], skip_special_tokens=True).strip(), int(encoded.input_ids.shape[1]), round(time.time() - started, 3)
return decode_generated_response(tokenizer, new_ids[0]), int(encoded.input_ids.shape[1]), round(time.time() - started, 3)


def parse_max_memory(items: list[str] | None) -> dict[int | str, str] | None:
Expand Down
26 changes: 23 additions & 3 deletions eval/benchmarks/memqa/scripts/run_metis_memqa.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,10 @@


from eval.methods.shared.memqa_io import audit_query_payload, query_prompt_for_style # noqa: E402
from metis.generation_utils import ( # noqa: E402
decode_generated_response,
resolve_generation_special_token_ids,
)

from eval.methods.shared.metis_loader import (
_chat_template,
Expand Down Expand Up @@ -95,17 +99,21 @@ def generate_answer(model: Any, tokenizer: Any, question: str, device: str, max_
prompt = query_prompt_for_style(question, query_style)
started = time.time()
prefix_ids = render_messages(tokenizer, [{"role": "user", "content": prompt}], add_generation_prompt=True, device=device)
eos_token_id, pad_token_id = resolve_generation_special_token_ids(
model,
tokenizer,
)
output_ids = model.generate(
input_ids=prefix_ids,
attention_mask=torch.ones_like(prefix_ids),
max_new_tokens=max_new_tokens,
do_sample=False,
use_cache=True,
pad_token_id=tokenizer.pad_token_id or tokenizer.eos_token_id,
eos_token_id=tokenizer.eos_token_id,
pad_token_id=pad_token_id,
eos_token_id=eos_token_id,
)
new_ids = output_ids[:, prefix_ids.shape[1] :]
generated = tokenizer.decode(new_ids[0], skip_special_tokens=True).strip()
generated = decode_generated_response(tokenizer, new_ids[0])
return generated, int(prefix_ids.shape[1]), round(time.time() - started, 3), prompt


Expand Down Expand Up @@ -136,6 +144,7 @@ def run(args: argparse.Namespace) -> Path:
ckpt_path,
device=args.device,
dtype=parse_dtype(args.dtype),
model_path=args.model_path or None,
device_map=args.device_map,
model_parallel_devices=args.model_parallel_devices,
max_memory=parse_max_memory(args.max_memory),
Expand All @@ -157,6 +166,11 @@ def run(args: argparse.Namespace) -> Path:
"baseline": "metis",
"model_label": args.model_label,
"model_path": str(ckpt_path),
"base_model_override": (
str(Path(args.model_path).expanduser().resolve())
if args.model_path
else None
),
"checkpoint_config": config_summary(ckpt_path),
"load_report": load_report,
"device": args.device,
Expand Down Expand Up @@ -209,6 +223,11 @@ def run(args: argparse.Namespace) -> Path:
"baseline": "metis",
"model_label": args.model_label,
"model_path": str(ckpt_path),
"base_model_override": (
str(Path(args.model_path).expanduser().resolve())
if args.model_path
else None
),
"device": args.device,
"input_device": str(input_device),
"physical_gpu": os.environ.get("CUDA_VISIBLE_DEVICES"),
Expand Down Expand Up @@ -253,6 +272,7 @@ def run(args: argparse.Namespace) -> Path:
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--checkpoint", required=True)
parser.add_argument("--model-path", default="")
parser.add_argument("--model-label", required=True)
parser.add_argument("--instances", type=Path, required=True)
parser.add_argument("--output", type=Path, required=True)
Expand Down
3 changes: 3 additions & 0 deletions eval/benchmarks/ood/scripts/run_memory_only.py
Original file line number Diff line number Diff line change
Expand Up @@ -94,6 +94,7 @@ def build_config(args: argparse.Namespace, instances_sha256: str, count: int) ->
common.update(
{
"checkpoint": args.model_path,
"base_model_path": args.base_model_path or None,
"device_map": args.device_map,
"model_parallel_devices": args.model_parallel_devices,
}
Expand Down Expand Up @@ -210,6 +211,7 @@ def load_runtime(args: argparse.Namespace) -> tuple[Any, dict[str, Any]]:
checkpoint,
device=args.device,
dtype=parse_dtype(args.dtype),
model_path=args.base_model_path or None,
device_map=args.device_map,
model_parallel_devices=args.model_parallel_devices,
max_memory=parse_max_memory(args.max_memory),
Expand Down Expand Up @@ -337,6 +339,7 @@ def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--method", required=True, choices=("metis", "delta_mem", "temp_lora"))
parser.add_argument("--model-path", required=True)
parser.add_argument("--base-model-path", default="")
parser.add_argument("--model-label", required=True)
parser.add_argument("--adapter-dir", default="")
parser.add_argument("--instances", type=Path, required=True)
Expand Down
2 changes: 1 addition & 1 deletion eval/environments/paper-eval-minimal-cu118.yml
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@ dependencies:
- --extra-index-url https://download.pytorch.org/whl/cu118
- torch==2.7.1+cu118
- numpy==1.26.4
- transformers==5.4.0
- transformers==5.10.4
- accelerate==1.13.0
- safetensors==0.7.0
- peft==0.19.1
Expand Down
5 changes: 4 additions & 1 deletion eval/experiments/main_tables/run.py
Original file line number Diff line number Diff line change
Expand Up @@ -109,11 +109,14 @@ def module_and_args(args: argparse.Namespace, instances: Path) -> tuple[str, lis
]
if task == "memqa":
if args.method == "metis":
return "eval.benchmarks.memqa.scripts.run_metis_memqa", [
method_args = [
"--checkpoint", checkpoint, "--model-label", args.model_label,
"--output", str(output), "--query-style", args.query_style,
"--max-new-tokens", str(args.max_new_tokens), *metis_loading, *common,
]
if args.model:
method_args += ["--model-path", args.model]
return "eval.benchmarks.memqa.scripts.run_metis_memqa", method_args
module = f"eval.methods.{args.method}.scripts.run_memqa_{args.method}"
method_args = ["--model-path", model, "--model-label", args.model_label, "--output", str(output)]
if args.method == "delta_mem":
Expand Down
7 changes: 5 additions & 2 deletions eval/experiments/ood/run.py
Original file line number Diff line number Diff line change
Expand Up @@ -73,8 +73,11 @@ def cell_commands(args: argparse.Namespace, dataset: str) -> list[list[str]]:
]
if args.method in {"metis", "temp_lora"}:
command.extend(["--device-map", args.device_map])
if args.method == "metis" and args.model_parallel_devices:
command.extend(["--model-parallel-devices", args.model_parallel_devices])
if args.method == "metis":
if args.model:
command.extend(["--base-model-path", args.model])
if args.model_parallel_devices:
command.extend(["--model-parallel-devices", args.model_parallel_devices])
for item in args.max_memory:
command.extend(["--max-memory", item])
if args.method == "delta_mem":
Expand Down
Loading