diff --git a/eval/benchmarks/memops/scripts/run_memop_memory_baseline.py b/eval/benchmarks/memops/scripts/run_memop_memory_baseline.py index de3eb9b..e40240f 100755 --- a/eval/benchmarks/memops/scripts/run_memop_memory_baseline.py +++ b/eval/benchmarks/memops/scripts/run_memop_memory_baseline.py @@ -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, @@ -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), @@ -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={ @@ -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", @@ -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, @@ -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, @@ -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), diff --git a/eval/benchmarks/memops/scripts/run_qwen_plain_context.py b/eval/benchmarks/memops/scripts/run_qwen_plain_context.py index 10bcf94..47d58fb 100755 --- a/eval/benchmarks/memops/scripts/run_qwen_plain_context.py +++ b/eval/benchmarks/memops/scripts/run_qwen_plain_context.py @@ -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() @@ -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, diff --git a/eval/benchmarks/memqa/scripts/run_base_context.py b/eval/benchmarks/memqa/scripts/run_base_context.py index 10dd30f..22072b6 100755 --- a/eval/benchmarks/memqa/scripts/run_base_context.py +++ b/eval/benchmarks/memqa/scripts/run_base_context.py @@ -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() @@ -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: diff --git a/eval/benchmarks/memqa/scripts/run_metis_memqa.py b/eval/benchmarks/memqa/scripts/run_metis_memqa.py index 5741250..4bc9287 100755 --- a/eval/benchmarks/memqa/scripts/run_metis_memqa.py +++ b/eval/benchmarks/memqa/scripts/run_metis_memqa.py @@ -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, @@ -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 @@ -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), @@ -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, @@ -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"), @@ -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) diff --git a/eval/benchmarks/ood/scripts/run_memory_only.py b/eval/benchmarks/ood/scripts/run_memory_only.py index 0856209..9dcf658 100644 --- a/eval/benchmarks/ood/scripts/run_memory_only.py +++ b/eval/benchmarks/ood/scripts/run_memory_only.py @@ -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, } @@ -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), @@ -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) diff --git a/eval/environments/paper-eval-minimal-cu118.yml b/eval/environments/paper-eval-minimal-cu118.yml index 39f51be..10dadb3 100644 --- a/eval/environments/paper-eval-minimal-cu118.yml +++ b/eval/environments/paper-eval-minimal-cu118.yml @@ -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 diff --git a/eval/experiments/main_tables/run.py b/eval/experiments/main_tables/run.py index 80e13ca..6a65d4c 100644 --- a/eval/experiments/main_tables/run.py +++ b/eval/experiments/main_tables/run.py @@ -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": diff --git a/eval/experiments/ood/run.py b/eval/experiments/ood/run.py index 395de11..5686695 100644 --- a/eval/experiments/ood/run.py +++ b/eval/experiments/ood/run.py @@ -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": diff --git a/eval/methods/shared/metis_loader.py b/eval/methods/shared/metis_loader.py index d9b1e50..cbe146b 100644 --- a/eval/methods/shared/metis_loader.py +++ b/eval/methods/shared/metis_loader.py @@ -30,6 +30,7 @@ from metis.configuration_metis import MetisConfig # noqa: E402 from metis.modeling_metis import MetisForCausalLM # noqa: E402 +from metis.weight_utils import _gemma4_text_model, _load_gemma4_pretrained # noqa: E402 DEPRECATED_CHECKPOINT_ALIASES = { @@ -113,6 +114,8 @@ def _normalize_backbone_type_for_local_wrappers(config: MetisConfig) -> tuple[st aliases = { "qwen3": "Qwen3", "qwen3_5": "Qwen3_5", + "gemma4": "Gemma4", + "llama": "Llama", } if isinstance(original, str) and original in aliases: backbone_meta["backbone_type"] = aliases[original] @@ -146,7 +149,20 @@ def _candidate_base_paths(raw_path: str | None) -> list[Path]: return candidates -def _base_model_path_from_checkpoint(ckpt_dir: Path, config: MetisConfig) -> str | Path: +def _base_model_path_from_checkpoint( + ckpt_dir: Path, + config: MetisConfig, + *, + model_path: str | Path | None = None, +) -> str | Path: + if model_path is not None: + explicit_path = Path(model_path).expanduser() + if not explicit_path.is_dir(): + raise FileNotFoundError( + f"Explicit base model override is not a directory: {explicit_path}" + ) + return explicit_path.resolve() + raw_paths: list[str] = [] manifest_path = ckpt_dir / "metis_delta_manifest.json" if manifest_path.exists(): @@ -184,17 +200,40 @@ def _load_backbone_weights_into_metis( base_model_path: str | Path, dtype: torch.dtype, ) -> dict[str, Any]: - backbone = AutoModelForCausalLM.from_pretrained( - str(base_model_path), - trust_remote_code=True, - dtype=dtype, + backbone_type = str( + (getattr(model.config, "backbone_meta", None) or {}).get( + "backbone_type", "" + ) ) + if backbone_type == "Gemma4": + backbone = _load_gemma4_pretrained( + str(base_model_path), + model.config, + dtype, + ) + text_model = _gemma4_text_model(backbone) + output_embeddings = backbone.get_output_embeddings() + if output_embeddings is None: + output_embeddings = text_model.embed_tokens + else: + backbone = AutoModelForCausalLM.from_pretrained( + str(base_model_path), + trust_remote_code=True, + dtype=dtype, + ) + text_model = backbone.model + output_embeddings = backbone.lm_head + metis_backbone = model.model.metis_backbone - metis_backbone.model.load_state_dict(backbone.model.state_dict(), strict=True) - metis_backbone.lm_head.load_state_dict(backbone.lm_head.state_dict(), strict=True) + metis_backbone.model.load_state_dict(text_model.state_dict(), strict=True) + metis_backbone.lm_head.load_state_dict( + output_embeddings.state_dict(), + strict=True, + ) report = { "base_model_path": str(base_model_path), "base_model_class": backbone.__class__.__name__, + "base_text_model_class": text_model.__class__.__name__, } del backbone gc.collect() @@ -235,11 +274,59 @@ def infer_metis_model_family(ckpt_dir: str | Path, config: MetisConfig) -> str: """Infer the Metis size family used for device policy checks.""" path_text = str(ckpt_dir).lower() + text_cfg = _text_config(config) + model_type = str(getattr(text_cfg, "model_type", "") or "").lower() + backbone_type = str( + (getattr(config, "backbone_meta", None) or {}).get( + "backbone_type", "" + ) + ).lower() + if model_type == "llama" or backbone_type == "llama": + name_candidates = [ + path_text, + str(getattr(config.backbone_configs, "_name_or_path", "") or "").lower(), + str( + (getattr(config, "backbone_meta", None) or {}).get( + "backbone_path", "" + ) + ).lower(), + ] + for candidate in name_candidates: + size_match = re.search( + r"(?= 8192 and num_layers >= 80: + return "llama_70b" + return "llama_unknown" + + if model_type.startswith("gemma4"): + name_candidates = [ + path_text, + str(getattr(config.backbone_configs, "_name_or_path", "") or "").lower(), + str( + (getattr(config, "backbone_meta", None) or {}).get( + "backbone_path", "" + ) + ).lower(), + ] + for candidate in name_candidates: + size_match = re.search( + r"(? str: - """Enforce the paper evaluation policy: 4B/9B single GPU, 27B two GPU.""" + """Enforce the released Qwen paper-evaluation device policy.""" family = infer_metis_model_family(ckpt_dir, config) devices = model_parallel_devices or [] - if device_map == "paired_layers" and len(devices) != 2: - raise ValueError( - "Metis evaluation device policy requires paired_layers to use exactly two devices; " - f"got {devices!r}. Use --model-parallel-devices cuda:0,cuda:1 for 27B, " - "and --device-map single for 4B/9B." - ) + if device_map == "paired_layers": + required_devices = 4 if family == "llama_70b" else 2 + if len(devices) != required_devices or len(set(devices)) != required_devices: + raise ValueError( + f"Metis evaluation device policy requires {family} paired_layers " + f"to use exactly {required_devices} distinct devices; got {devices!r}." + ) if family in {"4b", "9b"} and device_map != "single": raise ValueError( f"Metis evaluation device policy requires Metis {family.upper()} to run single-GPU; " @@ -279,6 +367,12 @@ def enforce_metis_device_policy( "Metis evaluation device policy requires Metis 27B to run on exactly two GPUs with " "--device-map paired_layers --model-parallel-devices cuda:0,cuda:1." ) + if family == "llama_70b" and device_map != "paired_layers": + raise ValueError( + "Metis evaluation device policy requires Llama 70B to run on exactly four GPUs with " + "--device-map paired_layers " + "--model-parallel-devices cuda:0,cuda:1,cuda:2,cuda:3." + ) return family @@ -306,6 +400,20 @@ def _stringify_device_map(device_map: Any) -> Any: return device_map +def _module_execution_device(module: Any, fallback: Any) -> Any: + hook_device = getattr( + getattr(module, "_hf_hook", None), + "execution_device", + None, + ) + if hook_device is not None: + return hook_device + try: + return next(module.parameters()).device + except StopIteration: + return fallback + + def _patch_model_parallel_commit_memory(model: MetisForCausalLM) -> None: def _commit_memory(self: MetisForCausalLM, outputs: Any, attention_mask: torch.Tensor | None = None) -> None: offset = self.config.memory_configs.get("commit_hidden_offset", 0) @@ -317,8 +425,17 @@ def _commit_memory(self: MetisForCausalLM, outputs: Any, attention_mask: torch.T continue layer_h = all_hidden[k + offset] layer_attention_mask = attention_mask - if layer_attention_mask is not None and layer_attention_mask.device != layer_h.device: - layer_attention_mask = layer_attention_mask.to(layer_h.device) + execution_device = _module_execution_device( + layer.hyper_memory, + layer_h.device, + ) + if ( + layer_attention_mask is not None + and layer_attention_mask.device != execution_device + ): + layer_attention_mask = layer_attention_mask.to( + execution_device + ) layer.hyper_memory.update_local_memory( layer_h, layer.local_memory, @@ -363,7 +480,13 @@ def _dispatch_model( from accelerate import dispatch_model, infer_auto_device_map model.to(dtype=dtype) - no_split = ["MetisBlock", "Qwen3_5DecoderLayer", "Qwen3DecoderLayer"] + no_split = [ + "MetisBlock", + "Qwen3_5DecoderLayer", + "Qwen3DecoderLayer", + "Gemma4DecoderLayer", + "Gemma4UnifiedDecoderLayer", + ] if device_map == "paired_layers": devices = paired_devices or parse_model_parallel_devices(model_parallel_devices, device) resolved_map = build_paired_layer_device_map(config, devices) @@ -420,6 +543,7 @@ def load_v2_full_checkpoint( device: str = "cuda:0", dtype: torch.dtype = torch.bfloat16, *, + model_path: str | Path | None = None, device_map: str = "single", model_parallel_devices: str | None = None, max_memory: dict[int | str, str] | None = None, @@ -432,7 +556,11 @@ def load_v2_full_checkpoint( state_path, checkpoint_format = _checkpoint_state_path(ckpt_dir) base_report: dict[str, Any] = {} if checkpoint_format == "delta": - base_path = _base_model_path_from_checkpoint(ckpt_dir, config) + base_path = _base_model_path_from_checkpoint( + ckpt_dir, + config, + model_path=model_path, + ) base_report = _load_backbone_weights_into_metis(model, base_path, dtype) state = load_file(state_path, device="cpu") @@ -471,6 +599,11 @@ def load_v2_full_checkpoint( or "query_norm" in key ) ] + if important_missing: + raise RuntimeError( + f"Missing required Metis tensors while loading {checkpoint_format} " + f"checkpoint {ckpt_dir}: {important_missing[:50]}" + ) model, dispatch_report = _dispatch_model( model, diff --git a/metis/backbone_wrappers/Gemma4_wrapper.py b/metis/backbone_wrappers/Gemma4_wrapper.py new file mode 100644 index 0000000..c716f08 --- /dev/null +++ b/metis/backbone_wrappers/Gemma4_wrapper.py @@ -0,0 +1,394 @@ +"""Metis wrapper for Gemma 4 and Gemma 4 Unified text backbones.""" + +import importlib +from collections import UserDict +from collections.abc import Callable +from functools import lru_cache + +import torch +import torch.nn as nn +from torch.utils.checkpoint import checkpoint as _ckpt +from transformers import Cache, DynamicCache, __version__ as transformers_version +from transformers.masking_utils import ( + create_causal_mask, + create_sliding_window_causal_mask, +) +from transformers.modeling_outputs import BaseModelOutputWithPast +from transformers.modeling_utils import ALL_ATTENTION_FUNCTIONS +from transformers.utils import TransformersKwargs +from transformers.utils.generic import merge_with_config_defaults +from transformers.utils.output_capturing import capture_outputs +from typing_extensions import Unpack + +from ..utils import CausalLMWrapperForMetis, DecoderLayerWrapperForMetis + + +def _is_unified_text_config(text_config) -> bool: + return str(getattr(text_config, "model_type", "")).startswith("gemma4_unified") + + +@lru_cache(maxsize=2) +def _load_gemma4_components(unified: bool): + """Load the matching HF implementation without making both variants mandatory.""" + if unified: + module_name = "transformers.models.gemma4_unified.modeling_gemma4_unified" + class_prefix = "Gemma4Unified" + minimum_version = "5.10.1" + else: + module_name = "transformers.models.gemma4.modeling_gemma4" + class_prefix = "Gemma4" + minimum_version = "5.5.0" + + try: + module = importlib.import_module(module_name) + except (ImportError, ModuleNotFoundError) as exc: + variant = "Gemma 4 Unified" if unified else "Gemma 4" + raise ImportError( + f"{variant} is unavailable in transformers {transformers_version}. " + f"Install transformers>={minimum_version} (5.10.1+ is recommended " + "when both Gemma 4 checkpoints are used)." + ) from exc + + return ( + getattr(module, f"{class_prefix}TextModel"), + getattr(module, f"{class_prefix}TextModelOutputWithPast"), + module.apply_rotary_pos_emb, + module.eager_attention_forward, + ) + + +class Gemma4DecoderLayerForMetis(DecoderLayerWrapperForMetis): + """Adapt one Hugging Face Gemma 4 text decoder layer to a Metis block.""" + + def __init__(self, config, raw_decoder): + super().__init__(config, raw_decoder) + text_config = getattr( + config.backbone_configs, "text_config", config.backbone_configs + ) + _, _, self._apply_rotary_pos_emb, self._eager_attention_forward = ( + _load_gemma4_components(_is_unified_text_config(text_config)) + ) + + def before_mixin( + self, + hidden_states: torch.Tensor, + attention_mask: torch.Tensor | None = None, + position_ids: torch.LongTensor | None = None, + past_key_values: Cache | None = None, + use_cache: bool | None = False, + cache_position: torch.LongTensor | None = None, + position_embeddings: tuple[torch.Tensor, torch.Tensor] | None = None, + shared_kv_states: dict[str, tuple[torch.Tensor, torch.Tensor]] | None = None, + per_layer_input: torch.Tensor | None = None, + **kwargs: Unpack[TransformersKwargs], + ): + del use_cache, per_layer_input + self_attn = self.raw_decoder.self_attn + # Set by NonMemoryMetisBlock: no memory read happens, skip the query copy. + skip_memory_query = getattr(self, "_skip_memory_query", False) + + def _attn(states): + residual = states + normed = self.raw_decoder.input_layernorm(states) + input_shape = normed.shape[:-1] + hidden_shape = (*input_shape, -1, self_attn.head_dim) + + query_states = self_attn.q_proj(normed).view(hidden_shape) + query_states = self_attn.q_norm(query_states) + + # Gemma normalizes Q before RoPE. Metis reads in that same + # position-independent query space. + query_for_memory = ( + None if skip_memory_query + else query_states.transpose(1, 2).contiguous().clone() + ) + + cos, sin = position_embeddings + query_states = self._apply_rotary_pos_emb( + query_states, cos, sin, unsqueeze_dim=2 + ).transpose(1, 2) + + if self_attn.is_kv_shared_layer: + if ( + shared_kv_states is None + or self_attn.layer_type not in shared_kv_states + ): + raise RuntimeError( + f"Missing shared KV states for Gemma 4 layer {self_attn.layer_idx}." + ) + key_states, value_states = shared_kv_states[self_attn.layer_type] + key_states = key_states.to(query_states.device) + value_states = value_states.to(query_states.device) + else: + key_states = self_attn.k_proj(normed).view(hidden_shape) + value_states = ( + self_attn.v_proj(normed).view(hidden_shape) + if self_attn.v_proj is not None + else key_states + ) + + key_states = self_attn.k_norm(key_states) + key_states = self._apply_rotary_pos_emb( + key_states, cos, sin, unsqueeze_dim=2 + ).transpose(1, 2) + value_states = self_attn.v_norm(value_states).transpose(1, 2) + + if past_key_values is not None and not self_attn.is_kv_shared_layer: + key_states, value_states = past_key_values.update( + key_states, value_states, self_attn.layer_idx + ) + if self_attn.store_full_length_kv: + if shared_kv_states is None: + raise RuntimeError("Gemma 4 shared_kv_states was not initialized.") + shared_kv_states[self_attn.layer_type] = key_states, value_states + + attention_interface: Callable = ALL_ATTENTION_FUNCTIONS.get_interface( + self_attn.config._attn_implementation, + self._eager_attention_forward, + ) + attn_output, _ = attention_interface( + self_attn, + query_states, + key_states, + value_states, + attention_mask, + dropout=self_attn.attention_dropout if self.training else 0.0, + scaling=self_attn.scaling, + sliding_window=self_attn.sliding_window, + position_ids=position_ids, + cache_position=cache_position, + **kwargs, + ) + + attn_output = attn_output.reshape(*input_shape, -1).contiguous() + attn_output = self_attn.o_proj(attn_output) + return query_for_memory, attn_output, residual + + if getattr(self, "_use_gradient_checkpointing", False) and self.training: + # DeepSpeed ZeRO-3 releases a module's gathered parameters back to + # zero-sized partitions after forward. With PyTorch 2.11, the + # non-reentrant checkpoint implementation then compares those + # released tensors against the full forward-time metadata and + # raises CheckpointError during recomputation. Reentrant + # checkpointing lets ZeRO-3's module hooks gather the parameters + # again as part of the backward recompute. + query_for_memory, attn_output, residual = _ckpt( + _attn, hidden_states, use_reentrant=True + ) + else: + query_for_memory, attn_output, residual = _attn(hidden_states) + + return ( + query_for_memory, + attn_output, + {"residual": residual}, + self_attn.o_proj, + ) + + def after_mixin( + self, + memory_carrier, + cache_dict, + hidden_states: torch.Tensor, + per_layer_input: torch.Tensor | None = None, + **kwargs: Unpack[TransformersKwargs], + ) -> torch.Tensor: + del hidden_states, kwargs + decoder = self.raw_decoder + hidden_states = cache_dict["residual"] + decoder.post_attention_layernorm( + memory_carrier + ) + + def _mlp(states, layer_input=None): + residual = states + states = decoder.pre_feedforward_layernorm(states) + states = decoder.mlp(states) + + if getattr(decoder, "enable_moe_block", False): + states_1 = decoder.post_feedforward_layernorm_1(states) + residual_flat = residual.reshape(-1, residual.shape[-1]) + _, top_k_weights, top_k_index = decoder.router(residual_flat) + states_2 = decoder.pre_feedforward_layernorm_2(residual_flat) + states_2 = decoder.experts(states_2, top_k_index, top_k_weights) + states_2 = states_2.reshape(residual.shape) + states_2 = decoder.post_feedforward_layernorm_2(states_2) + states = states_1 + states_2 + + states = decoder.post_feedforward_layernorm(states) + states = residual + states + + if getattr(decoder, "hidden_size_per_layer_input", 0): + if layer_input is None: + raise ValueError("Gemma 4 per-layer input is required by this config.") + residual = states + states = decoder.per_layer_input_gate(states) + states = decoder.act_fn(states) * layer_input + states = decoder.per_layer_projection(states) + states = decoder.post_per_layer_input_norm(states) + states = residual + states + + return states * decoder.layer_scalar + + if getattr(self, "_use_gradient_checkpointing", False) and self.training: + if per_layer_input is None: + hidden_states = _ckpt(_mlp, hidden_states, use_reentrant=True) + else: + hidden_states = _ckpt( + _mlp, hidden_states, per_layer_input, use_reentrant=True + ) + else: + hidden_states = _mlp(hidden_states, per_layer_input) + + return hidden_states + + +class Gemma4CausalLMForMetis(CausalLMWrapperForMetis): + """Text-only Gemma 4 shell whose decoder layers are executed by Metis.""" + + def __init__(self, config): + super().__init__(config) + text_config = getattr( + config.backbone_configs, "text_config", config.backbone_configs + ) + text_model_cls, output_cls, _, _ = _load_gemma4_components( + _is_unified_text_config(text_config) + ) + + self.model = text_model_cls(text_config) + self.vocab_size = text_config.vocab_size + self.lm_head = nn.Linear( + text_config.hidden_size, + text_config.vocab_size, + bias=False, + ) + self._output_cls = output_cls + self.final_logit_softcapping = getattr( + text_config, "final_logit_softcapping", None + ) + + def get_decoder_layer_by_id(self, layer_id: int): + return self.model.layers[layer_id] + + def project_logits(self, hidden_states: torch.Tensor) -> torch.Tensor: + logits = self.lm_head(hidden_states) + if self.final_logit_softcapping is not None: + logits = logits / self.final_logit_softcapping + logits = torch.tanh(logits) + logits = logits * self.final_logit_softcapping + return logits + + @merge_with_config_defaults + @capture_outputs + def forward_with_memory( + self, + input_ids: torch.LongTensor | None = None, + attention_mask: torch.Tensor | dict | None = None, + position_ids: torch.LongTensor | None = None, + past_key_values: Cache | None = None, + inputs_embeds: torch.FloatTensor | None = None, + per_layer_inputs: torch.Tensor | None = None, + use_cache: bool | None = None, + cache_position: torch.LongTensor | None = None, + shared_kv_states: dict[str, tuple[torch.Tensor, torch.Tensor]] | None = None, + **kwargs: Unpack[TransformersKwargs], + ) -> BaseModelOutputWithPast: + if (input_ids is None) ^ (inputs_embeds is not None): + raise ValueError("You must specify exactly one of input_ids or inputs_embeds") + if input_ids is not None and per_layer_inputs is not None: + raise ValueError("You cannot specify per_layer_inputs if input_ids is provided") + + original_input_ids = input_ids + if inputs_embeds is None: + inputs_embeds = self.model.embed_tokens(input_ids) + + if getattr(self.model, "hidden_size_per_layer_input", 0): + if per_layer_inputs is None: + per_layer_inputs = self.model.get_per_layer_inputs( + original_input_ids, inputs_embeds + ) + per_layer_inputs = self.model.project_per_layer_inputs( + inputs_embeds, per_layer_inputs + ) + + if use_cache and past_key_values is None: + past_key_values = DynamicCache(config=self.model.config) + + if position_ids is None: + if cache_position is not None: + position_ids = cache_position.unsqueeze(0) + else: + past_seen_tokens = ( + past_key_values.get_seq_length() + if past_key_values is not None + else 0 + ) + position_ids = torch.arange( + inputs_embeds.shape[1], device=inputs_embeds.device + ) + past_seen_tokens + position_ids = position_ids.unsqueeze(0) + + if not isinstance(causal_mask_mapping := attention_mask, dict): + mask_kwargs = { + "config": self.model.config, + "inputs_embeds": inputs_embeds, + "attention_mask": attention_mask, + "past_key_values": past_key_values, + "position_ids": position_ids, + } + causal_mask_mapping = { + "full_attention": create_causal_mask(**mask_kwargs), + "sliding_attention": create_sliding_window_causal_mask(**mask_kwargs), + } + + output_hidden_states = kwargs.get( + "output_hidden_states", self.config.output_hidden_states + ) + all_hidden_states = () if output_hidden_states else None + + hidden_states = inputs_embeds + position_embeddings = { + layer_type: self.model.rotary_emb( + hidden_states, position_ids, layer_type + ) + for layer_type in self.model.unique_layer_types + } + if shared_kv_states is None: + shared_kv_states = UserDict() + + for layer_idx in range(self.model.config.num_hidden_layers): + if output_hidden_states: + all_hidden_states += (hidden_states,) + + layer_type = self.model.config.layer_types[layer_idx] + per_layer_input = ( + per_layer_inputs[:, :, layer_idx, :] + if per_layer_inputs is not None + else None + ) + hidden_states = self._metis_blocks_ref[layer_idx]( + hidden_states, + per_layer_input=per_layer_input, + shared_kv_states=shared_kv_states, + position_embeddings=position_embeddings[layer_type], + attention_mask=causal_mask_mapping[layer_type], + position_ids=position_ids, + past_key_values=past_key_values, + use_cache=use_cache, + cache_position=cache_position, + **kwargs, + ) + + hidden_states = self.model.norm(hidden_states) + if output_hidden_states: + all_hidden_states += (hidden_states,) + + return self._output_cls( + last_hidden_state=hidden_states, + past_key_values=past_key_values, + hidden_states=all_hidden_states, + shared_kv_states=( + shared_kv_states + if kwargs.get("return_shared_kv_states", False) + else None + ), + ) diff --git a/metis/backbone_wrappers/Llama_wrapper.py b/metis/backbone_wrappers/Llama_wrapper.py index 8f29d6a..d4fb2b8 100644 --- a/metis/backbone_wrappers/Llama_wrapper.py +++ b/metis/backbone_wrappers/Llama_wrapper.py @@ -20,6 +20,30 @@ from ..utils import CausalLMWrapperForMetis, DecoderLayerWrapperForMetis +def _create_llama_causal_mask( + *, + config, + inputs_embeds, + attention_mask, + past_key_values, + position_ids, +): + """Call the Transformers 5.x causal-mask API without a removed keyword. + + Transformers 5.4 accepted a deprecated ``cache_position`` keyword but did + not use it. Transformers 5.10 removed that keyword. Omitting it preserves + the mask semantics and keeps the wrapper compatible with both releases. + """ + + return create_causal_mask( + config=config, + inputs_embeds=inputs_embeds, + attention_mask=attention_mask, + past_key_values=past_key_values, + position_ids=position_ids, + ) + + class LlamaDecoderLayerForMetis(DecoderLayerWrapperForMetis): def __init__(self, config, raw_decoder): super().__init__(config, raw_decoder) @@ -163,11 +187,10 @@ def forward_with_memory( if position_ids is None: position_ids = cache_position.unsqueeze(0) - causal_mask = create_causal_mask( + causal_mask = _create_llama_causal_mask( config=self.model.config, inputs_embeds=inputs_embeds, attention_mask=attention_mask, - cache_position=cache_position, past_key_values=past_key_values, position_ids=position_ids, ) diff --git a/metis/backbone_wrappers/Qwen3_5_wrapper.py b/metis/backbone_wrappers/Qwen3_5_wrapper.py index 5aa5534..fee0ffe 100644 --- a/metis/backbone_wrappers/Qwen3_5_wrapper.py +++ b/metis/backbone_wrappers/Qwen3_5_wrapper.py @@ -43,6 +43,79 @@ def _zero3_safe_checkpoint(function, hidden_states): ) +def _create_qwen3_5_causal_mask( + *, + config, + inputs_embeds, + attention_mask, + past_key_values, + position_ids, +): + """Call the Transformers 5.x mask API shared by 5.4 and 5.10. + + Transformers 5.4 accepted a deprecated ``cache_position`` keyword but did + not use it. Transformers 5.10 removed that keyword. Omitting it preserves + 5.4 behavior and keeps the wrapper compatible with both releases. + """ + + return create_causal_mask( + config=config, + inputs_embeds=inputs_embeds, + attention_mask=attention_mask, + past_key_values=past_key_values, + position_ids=position_ids, + ) + + +def _is_compatible_qwen3_5_generation_cache(past_key_values, config): + if not isinstance(past_key_values, Qwen3_5DynamicCache): + return False + + cache_layers = getattr(past_key_values, "layers", None) + layer_types = getattr(config, "layer_types", None) + if cache_layers is None or layer_types is None: + return True + if len(cache_layers) != len(layer_types): + return False + for layer_type, cache_layer in zip(layer_types, cache_layers): + if layer_type == "linear_attention": + required_names = ("conv_states", "recurrent_states") + elif layer_type == "full_attention": + required_names = ("keys", "values") + else: + continue + if not all(hasattr(cache_layer, name) for name in required_names): + return False + return True + + +def _prepare_qwen3_5_generation_cache( + past_key_values, + *, + config, + use_cache, +): + """Create the configured hybrid cache on the first generation call. + + Transformers generation may pre-create an empty generic cache. In 5.4 the + Qwen-specific cache class must replace it; in 5.10 ``DynamicCache`` uses + the model config to build the full/linear hybrid layers. A foreign cache + must be rebuilt even when populated; a compatible populated cache is + preserved on later decode steps. + """ + + if not use_cache: + return past_key_values + if past_key_values is None: + return Qwen3_5DynamicCache(config=config) + if not _is_compatible_qwen3_5_generation_cache(past_key_values, config): + # A populated foreign cache cannot safely represent Qwen3.5's hybrid state. + return Qwen3_5DynamicCache(config=config) + if past_key_values.get_seq_length() == 0: + return Qwen3_5DynamicCache(config=config) + return past_key_values + + class Qwen3_5DecoderLayerForMetis(DecoderLayerWrapperForMetis): def __init__(self, config, raw_decoder): super().__init__(config, raw_decoder) @@ -256,13 +329,11 @@ def forward_with_memory( if inputs_embeds is None: inputs_embeds = self.model.embed_tokens(input_ids) - # generate() in transformers 5.x pre-creates a standard DynamicCache. - # Qwen3.5's hybrid (full + linear attention) model requires Qwen3_5DynamicCache, - # which manages recurrent states for linear attention layers. - # On the first call the incoming cache is always empty, so replacement is safe. - # On subsequent calls the cache is already Qwen3_5DynamicCache and is left as-is. - if use_cache and not isinstance(past_key_values, Qwen3_5DynamicCache): - past_key_values = Qwen3_5DynamicCache(config=self.model.config) + past_key_values = _prepare_qwen3_5_generation_cache( + past_key_values, + config=self.model.config, + use_cache=use_cache, + ) if cache_position is None: past_seen_tokens = past_key_values.get_seq_length() if past_key_values is not None else 0 @@ -282,11 +353,10 @@ def forward_with_memory( else: text_position_ids = None - causal_mask = create_causal_mask( + causal_mask = _create_qwen3_5_causal_mask( config=self.config, inputs_embeds=inputs_embeds, attention_mask=attention_mask, - cache_position=cache_position, past_key_values=past_key_values, position_ids=text_position_ids, ) diff --git a/metis/checkpoint_utils.py b/metis/checkpoint_utils.py index dfa3fa7..8c933b1 100644 --- a/metis/checkpoint_utils.py +++ b/metis/checkpoint_utils.py @@ -242,6 +242,7 @@ def _build_metis_from_delta_config( metis_hyper_memory_type=_memory_arg(memory, "metis_hyper_memory_type", "StraightThroughAlphaTopPGatedDeltaRuleMetisHyperMemory"), metis_local_memory_type=_memory_arg(memory, "metis_local_memory_type", "NormalizedDeltaNetMetisLocalMemory"), update_ratio=float(_memory_arg(memory, "update_ratio", 0.9)), + forget_ratio=float(_memory_arg(memory, "forget_ratio", 1.0)), commit_hidden_offset=int(_memory_arg(memory, "commit_hidden_offset", 0)), mem_norm_init=float(_memory_arg(memory, "mem_norm_init", 1.0)), uniform_num_selected=int(_memory_arg(memory, "uniform_num_selected", 16)), @@ -256,6 +257,9 @@ def _build_metis_from_delta_config( gated_delta_beta_init=float(_memory_arg(memory, "gated_delta_beta_init", 1.0)), qk_kernel_type=_memory_arg(memory, "qk_kernel_type", "elu_plus_one"), metis_reweight_gamma=float(_memory_arg(memory, "metis_reweight_gamma", 0.9)), + memory_on_sliding_layers=bool( + _memory_arg(memory, "memory_on_sliding_layers", True) + ), ) model.config.backbone_meta = { **backbone_meta, @@ -335,12 +339,47 @@ def load_metis_delta_into_model( return delta = load_file(str(checkpoint / DELTA_WEIGHTS_NAME), device="cpu") + manifest = load_delta_manifest(checkpoint) + manifest_tensors = manifest.get("tensors") or {} + manifest_keys = set(manifest_tensors) + delta_keys = set(delta) + missing_from_file = sorted(manifest_keys - delta_keys) + absent_from_manifest = sorted(delta_keys - manifest_keys) + if missing_from_file or absent_from_manifest: + raise RuntimeError( + f"Delta manifest/tensor mismatch for {checkpoint}: " + f"missing_from_file={missing_from_file[:10]} " + f"absent_from_manifest={absent_from_manifest[:10]}" + ) + for key, tensor in delta.items(): + declared = manifest_tensors[key] + declared_shape = list(declared.get("shape") or []) + declared_dtype = str(declared.get("dtype") or "") + actual_dtype = _safe_dtype_name(tensor) + if declared_shape != list(tensor.shape) or declared_dtype != actual_dtype: + raise RuntimeError( + f"Delta manifest metadata mismatch for {key!r}: " + f"declared shape/dtype={declared_shape}/{declared_dtype}, " + f"actual={list(tensor.shape)}/{actual_dtype}" + ) incompatible = unwrapped.load_state_dict(delta, strict=False) if incompatible.unexpected_keys: preview = ", ".join(incompatible.unexpected_keys[:10]) raise RuntimeError( f"Delta checkpoint has {len(incompatible.unexpected_keys)} unexpected keys: {preview}" ) + missing_metis = [ + key + for key in incompatible.missing_keys + if key.startswith("model.metis_blocks.") + ] + if missing_metis: + preview = ", ".join(missing_metis[:10]) + raise RuntimeError( + f"Delta checkpoint is missing {len(missing_metis)} Metis block " + f"tensors present in the rebuilt model: {preview}. Checkpoint " + "memory_configs and rebuilt architecture do not match." + ) logger.info("Loaded Metis delta weights from %s (%d tensors)", checkpoint, len(delta)) diff --git a/metis/configuration_metis.py b/metis/configuration_metis.py index 3053981..e2b1327 100644 --- a/metis/configuration_metis.py +++ b/metis/configuration_metis.py @@ -1,6 +1,13 @@ from transformers import PretrainedConfig, AutoConfig +def _gemma4_config_error() -> ImportError: + return ImportError( + "Gemma 4 support requires transformers>=5.10.1 so both gemma4 and " + "gemma4_unified configs are registered." + ) + + class MetisConfig(PretrainedConfig): model_type = "metis" @@ -23,7 +30,12 @@ def __init__( if isinstance(backbone_configs, dict) and "model_type" in backbone_configs: cfg = dict(backbone_configs) model_type = cfg.pop("model_type") - backbone_configs = AutoConfig.for_model(model_type, **cfg) + try: + backbone_configs = AutoConfig.for_model(model_type, **cfg) + except ValueError as exc: + if str(model_type).startswith("gemma4"): + raise _gemma4_config_error() from exc + raise self.backbone_configs = backbone_configs else: # Eagerly load only when an explicit path is provided. @@ -31,7 +43,15 @@ def __init__( # diff/serialization), backbone_path is "" so this is skipped. backbone_path = self.backbone_meta.get('backbone_path', '') if backbone_path: - self.backbone_configs = AutoConfig.from_pretrained(backbone_path) + try: + self.backbone_configs = AutoConfig.from_pretrained(backbone_path) + except ValueError as exc: + backbone_type = str( + self.backbone_meta.get("backbone_type", "") + ).lower() + if backbone_type == "gemma4": + raise _gemma4_config_error() from exc + raise else: self.backbone_configs = None @@ -53,6 +73,11 @@ def __init__( self.num_hidden_layers = text_cfg.num_hidden_layers self.num_attention_heads = text_cfg.num_attention_heads self.hidden_size = text_cfg.hidden_size - self.bos_token_id = text_cfg.bos_token_id - self.eos_token_id = text_cfg.eos_token_id - self.pad_token_id = text_cfg.pad_token_id + # Composite instruction checkpoints may keep generation stop IDs + # on the outer config while their text config retains only the + # base text EOS ID. + for token_id_name in ("bos_token_id", "eos_token_id", "pad_token_id"): + token_id = getattr(self.backbone_configs, token_id_name, None) + if token_id is None: + token_id = getattr(text_cfg, token_id_name) + setattr(self, token_id_name, token_id) diff --git a/metis/dev_beta/metis_block.py b/metis/dev_beta/metis_block.py index 3327e25..06c1cd5 100644 --- a/metis/dev_beta/metis_block.py +++ b/metis/dev_beta/metis_block.py @@ -1,3 +1,5 @@ +import copy + from transformers import GradientCheckpointingLayer import torch import torch.nn as nn @@ -7,8 +9,28 @@ from abc import ABC +def _layer_type_of(config, layer_idx, raw_decoder) -> str: + """Resolve the attention type used by one instantiated decoder layer.""" + text_cfg = getattr(config.backbone_configs, "text_config", config.backbone_configs) + layer_types = getattr(text_cfg, "layer_types", None) + if layer_types is not None and layer_idx < len(layer_types): + return layer_types[layer_idx] + self_attn = getattr(raw_decoder, "self_attn", None) + return getattr( + raw_decoder, + "layer_type", + getattr(self_attn, "layer_type", "full_attention"), + ) + + def create_metis_block(config, layer_idx, raw_decoder): - if getattr(raw_decoder, "layer_type", "full_attention") == "linear_attention": + layer_type = _layer_type_of(config, layer_idx, raw_decoder) + if layer_type == "linear_attention": + return NonMemoryMetisBlock(config, layer_idx, raw_decoder) + if ( + layer_type == "sliding_attention" + and not config.memory_configs.get("memory_on_sliding_layers", True) + ): return NonMemoryMetisBlock(config, layer_idx, raw_decoder) return eval(config.memory_configs['metis_block_type'])(config, layer_idx, raw_decoder) @@ -19,12 +41,58 @@ def _is_boundary_memory_layer(config, layer_idx: int) -> bool: layer_types = getattr(text_cfg, "layer_types", None) if layer_types is None: return layer_idx in {0, num_layers - 1} - memory_layers = [ - i for i, lt in enumerate(layer_types) - if lt != "linear_attention" - ] + skip_types = {"linear_attention"} + if not config.memory_configs.get("memory_on_sliding_layers", True): + skip_types.add("sliding_attention") + memory_layers = [i for i, layer_type in enumerate(layer_types) if layer_type not in skip_types] return bool(memory_layers) and layer_idx in {memory_layers[0], memory_layers[-1]} + +def _attention_geometry(config, raw_decoder): + """Return the effective query/KV head counts and head size for one layer.""" + text_cfg = getattr(config.backbone_configs, "text_config", config.backbone_configs) + self_attn = getattr(raw_decoder, "self_attn", None) + num_q_heads = int(text_cfg.num_attention_heads) + default_head_dim = getattr( + text_cfg, "head_dim", text_cfg.hidden_size // num_q_heads + ) + head_dim = int(getattr(self_attn, "head_dim", default_head_dim)) + k_proj = getattr(self_attn, "k_proj", None) + if k_proj is not None and getattr(k_proj, "out_features", 0) % head_dim == 0: + num_kv_heads = k_proj.out_features // head_dim + else: + num_kv_groups = int(getattr(self_attn, "num_key_value_groups", 1)) + num_kv_heads = num_q_heads // num_kv_groups + return num_q_heads, int(num_kv_heads), head_dim + + +def _config_for_attention_layer(config, raw_decoder): + """Return a shallow config view aligned with this layer's geometry.""" + num_q_heads, num_kv_heads, head_dim = _attention_geometry(config, raw_decoder) + text_cfg = getattr(config.backbone_configs, "text_config", config.backbone_configs) + current = ( + int(text_cfg.num_attention_heads), + int(getattr(text_cfg, "num_key_value_heads", num_q_heads)), + int(getattr(text_cfg, "head_dim", text_cfg.hidden_size // num_q_heads)), + ) + effective = (num_q_heads, num_kv_heads, head_dim) + if current == effective: + return config + + layer_config = copy.copy(config) + backbone_config = copy.copy(config.backbone_configs) + if hasattr(config.backbone_configs, "text_config"): + layer_text_config = copy.copy(text_cfg) + backbone_config.text_config = layer_text_config + else: + layer_text_config = backbone_config + layer_text_config.num_attention_heads = num_q_heads + layer_text_config.num_key_value_heads = num_kv_heads + layer_text_config.head_dim = head_dim + layer_config.backbone_configs = backbone_config + return layer_config + + class NonMemoryMetisBlock(GradientCheckpointingLayer): """Block with no memory modules — used for debugging / ablation. Passes through the backbone decoder with no memory read or write.""" @@ -40,6 +108,7 @@ def __init__(self, config, layer_idx, raw_decoder): self.last_attention_branch_norm = None self.last_memory_attention_norm_ratio = None self._backbone_decoder_ref = [create_metis_decoder_layer(config, raw_decoder)] + self._backbone_decoder_ref[0]._skip_memory_query = True @property def backbone_decoder(self): @@ -61,8 +130,9 @@ def __init__(self, config, layer_idx: int, raw_decoder): self.config = config self.layer_idx = layer_idx - self.local_memory = create_metis_local_memory(config) - self.hyper_memory = create_metis_hyper_memory(config) + layer_config = _config_for_attention_layer(config, raw_decoder) + self.local_memory = create_metis_local_memory(layer_config) + self.hyper_memory = create_metis_hyper_memory(layer_config) self._backbone_decoder_ref = [create_metis_decoder_layer(config, raw_decoder)] self.hyper_memory.register_raw_decoder(self._backbone_decoder_ref) @@ -101,8 +171,9 @@ def _record_query_norm(self, query_for_memory): def _init_learned_query(self, raw_decoder): text_cfg = getattr(self.config.backbone_configs, 'text_config', self.config.backbone_configs) hidden_dim = text_cfg.hidden_size - self.query_num_heads = text_cfg.num_attention_heads - self.query_head_dim = getattr(text_cfg, "head_dim", hidden_dim // self.query_num_heads) + self.query_num_heads, _, self.query_head_dim = _attention_geometry( + self.config, raw_decoder + ) mem_dim = self.query_num_heads * self.query_head_dim self_attn = getattr(raw_decoder, "self_attn", None) @@ -161,8 +232,8 @@ class NormedNaiveMetisBlock(NaiveMetisBlock): def __init__(self, config, layer_idx: int, raw_decoder): super().__init__(config, layer_idx, raw_decoder) - text_cfg = getattr(config.backbone_configs, 'text_config', config.backbone_configs) - mem_dim = text_cfg.num_attention_heads * text_cfg.head_dim + num_q_heads, _, head_dim = _attention_geometry(config, raw_decoder) + mem_dim = num_q_heads * head_dim ln = raw_decoder.input_layernorm norm_cls = type(ln) diff --git a/metis/generation_utils.py b/metis/generation_utils.py new file mode 100644 index 0000000..c16635c --- /dev/null +++ b/metis/generation_utils.py @@ -0,0 +1,92 @@ +"""Small generation-config helpers shared by Metis and base-model runners.""" + +from __future__ import annotations + +from typing import Any + + +def resolve_generation_special_token_ids( + model: Any, + tokenizer: Any, +) -> tuple[int | list[int], int]: + """Resolve model stops without dropping the tokenizer's chat boundary. + + Some Qwen checkpoints retain ``<|endoftext|>`` in the model generation + config while the tokenizer uses ``<|im_end|>`` to end an assistant turn. + Generation must stop on either token. Gemma's model-native multi-EOS list + and valid PAD token zero must remain unchanged. + """ + + generation_config = getattr(model, "generation_config", None) + model_eos_token_id = getattr(generation_config, "eos_token_id", None) + tokenizer_eos_token_id = getattr(tokenizer, "eos_token_id", None) + pad_token_id = getattr(generation_config, "pad_token_id", None) + + def as_list(value: Any) -> list[int]: + if value is None: + return [] + if isinstance(value, (list, tuple)): + return list(value) + return [value] + + eos_ids = as_list(model_eos_token_id) + for token_id in as_list(tokenizer_eos_token_id): + if token_id not in eos_ids: + eos_ids.append(token_id) + if not eos_ids: + raise ValueError("Cannot resolve an EOS token id") + eos_token_id: int | list[int] = ( + eos_ids + if isinstance(model_eos_token_id, (list, tuple)) or len(eos_ids) > 1 + else eos_ids[0] + ) + + if pad_token_id is None: + pad_token_id = tokenizer.pad_token_id + if pad_token_id is None: + if isinstance(eos_token_id, (list, tuple)): + if not eos_token_id: + raise ValueError("Cannot resolve pad token from an empty EOS list") + pad_token_id = eos_token_id[0] + else: + pad_token_id = eos_token_id + + return eos_token_id, pad_token_id + + +def decode_generated_response( + tokenizer: Any, + generated_ids: Any, +) -> str: + """Decode chat output and remove model-native reasoning channel markup.""" + + raw_response = tokenizer.decode( + generated_ids, skip_special_tokens=False + ) + response = None + parse_response = getattr(tokenizer, "parse_response", None) + if callable(parse_response): + try: + parsed = parse_response(raw_response) + except (AttributeError, TypeError, ValueError): + parsed = None + if isinstance(parsed, dict): + content = parsed.get("content") + if isinstance(content, str): + response = content + elif isinstance(parsed, str): + response = parsed + + if response is None: + response = tokenizer.decode( + generated_ids, skip_special_tokens=True + ) + + # Gemma 4 may greedily repeat the already-prefilled disabled-thinking + # channel tail as `thought\n` before its final answer. The + # tokenizer parser intentionally leaves that incomplete leading channel + # untouched, so remove only this exact structural prefix. + duplicate_disabled_thinking_prefix = "thought\n" + if response.startswith(duplicate_disabled_thinking_prefix): + response = response[len(duplicate_disabled_thinking_prefix) :] + return response.strip() diff --git a/metis/memory_utils.py b/metis/memory_utils.py index 11e9a69..57fe63f 100644 --- a/metis/memory_utils.py +++ b/metis/memory_utils.py @@ -31,7 +31,11 @@ def encode_and_commit_memory( hooks = [] for k, block in enumerate(model.model.metis_blocks): - if not is_full_attention(block): + if ( + block.local_memory is None + or block.hyper_memory is None + or not is_full_attention(block) + ): continue def _hook(mod, inp, out, idx=k): @@ -47,6 +51,10 @@ def _hook(mod, inp, out, idx=k): for k, block in enumerate(model.model.metis_blocks): if k in captured: + if block.local_memory is None or block.hyper_memory is None: + raise RuntimeError( + f"Captured layer {k} has no writable Metis memory." + ) block.hyper_memory.update_local_memory( captured[k], block.local_memory, diff --git a/metis/modeling_metis.py b/metis/modeling_metis.py index d047e83..b6cdbf2 100644 --- a/metis/modeling_metis.py +++ b/metis/modeling_metis.py @@ -99,6 +99,9 @@ def forward( attention_mask_1d: torch.Tensor | None = None, **kwargs, ) -> CausalLMOutputWithPast: + return_hidden_states = kwargs.get("output_hidden_states") + if return_hidden_states is None: + return_hidden_states = bool(self.config.output_hidden_states) if commit_memory: kwargs['output_hidden_states'] = True @@ -115,7 +118,9 @@ def forward( hidden_states = outputs.last_hidden_state slice_indices = slice(-logits_to_keep, None) if isinstance(logits_to_keep, int) else logits_to_keep - logits = self.model.metis_backbone.lm_head(hidden_states[:, slice_indices, :]) + logits = self.model.metis_backbone.project_logits( + hidden_states[:, slice_indices, :] + ) loss = None if labels is not None: @@ -137,6 +142,6 @@ def forward( loss=loss, logits=logits, past_key_values=outputs.past_key_values, - hidden_states=outputs.hidden_states, + hidden_states=outputs.hidden_states if return_hidden_states else None, attentions=outputs.attentions, ) diff --git a/metis/utils.py b/metis/utils.py index 03160e9..aa0ecaa 100644 --- a/metis/utils.py +++ b/metis/utils.py @@ -74,7 +74,10 @@ def register_metis_blocks(self, metis_blocks): def get_decoder_layer_by_id(self, layer_id: int): raise NotImplementedError + + def project_logits(self, hidden_states: torch.Tensor) -> torch.Tensor: + """Project hidden states, allowing backbone-specific post-processing.""" + return self.lm_head(hidden_states) def forward_with_memory(self, **kwargs): raise NotImplementedError - diff --git a/metis/weight_utils.py b/metis/weight_utils.py index 4180a60..ec395a8 100644 --- a/metis/weight_utils.py +++ b/metis/weight_utils.py @@ -16,6 +16,7 @@ "qwen3_5": "Qwen3_5", "qwen3": "Qwen3", "llama": "Llama", + "gemma4": "Gemma4", } @@ -329,6 +330,12 @@ def _init_hyper_memory_from_backbone(model: MetisForCausalLM) -> None: # full_attention: k_proj → W_k, v_proj → W_v src_k = getattr(self_attn, "k_proj", None) src_v = getattr(self_attn, "v_proj", None) + if ( + src_k is not None + and src_v is None + and getattr(self_attn, "use_alternative_attention", False) + ): + src_v = src_k if src_k is None or src_v is None: linear_attn = getattr(raw_decoder, "linear_attn", None) @@ -407,6 +414,42 @@ def _init_hyper_memory_from_backbone(model: MetisForCausalLM) -> None: print(f"[weight_utils] Hyper-memory initialized from backbone projections: {inited} layers (skipped {skipped}).") +def _gemma4_text_model(backbone): + """Locate the text decoder inside Gemma 4 or Gemma 4 Unified shells.""" + candidates = [ + getattr(getattr(backbone, "model", None), "language_model", None), + getattr(backbone, "language_model", None), + getattr(backbone, "model", None), + ] + for candidate in candidates: + if ( + candidate is not None + and hasattr(candidate, "layers") + and hasattr(candidate, "embed_tokens") + ): + return candidate + raise RuntimeError( + f"Could not locate the Gemma 4 text model inside {type(backbone).__name__}." + ) + + +def _load_gemma4_pretrained(backbone_path: str, config, dtype: torch.dtype): + """Load the matching Hugging Face Gemma shell.""" + model_type = str(getattr(config.backbone_configs, "model_type", "")) + if model_type.startswith("gemma4_unified"): + try: + from transformers import AutoModelForMultimodalLM + except ImportError as exc: + raise ImportError( + "Gemma 4 Unified requires transformers>=5.10.1." + ) from exc + return AutoModelForMultimodalLM.from_pretrained( + backbone_path, + dtype=dtype, + ) + return AutoModelForCausalLM.from_pretrained(backbone_path, dtype=dtype) + + def load_metis_from_backbone( backbone_path: str, backbone_type: str = "qwen3_5", @@ -416,6 +459,7 @@ def load_metis_from_backbone( metis_hyper_memory_type: str = "StraightThroughAlphaTopPGatedDeltaRuleMetisHyperMemory", metis_local_memory_type: str = "NormalizedDeltaNetMetisLocalMemory", update_ratio: float = 0.9, + forget_ratio: float = 1.0, commit_hidden_offset: int = 0, mem_norm_init: float = 1.0, uniform_num_selected: int = 16, @@ -430,18 +474,21 @@ def load_metis_from_backbone( gated_delta_beta_init: float = 1.0, qk_kernel_type: str = "elu_plus_one", metis_reweight_gamma: float = 0.9, + memory_on_sliding_layers: bool = True, ) -> tuple[MetisForCausalLM, "AutoTokenizer"]: """Build a MetisForCausalLM and load pretrained backbone weights into it. Args: backbone_path: Local path (or HF hub id) to the backbone model. - backbone_type: One of ``qwen3_5`` / ``qwen3`` / ``llama`` (lowercase). + backbone_type: One of ``qwen3_5`` / ``qwen3`` / ``llama`` / + ``gemma4`` (lowercase). device: Target device. dtype: Weight dtype (e.g. torch.float16 / torch.bfloat16). metis_block_type: Which MetisBlock variant to use. metis_hyper_memory_type: Which HyperMemory variant to use. metis_local_memory_type: Which LocalMemory variant to use. update_ratio: Blend factor for memory write (1.0 = full update). + forget_ratio: Stored for checkpoint/config compatibility. commit_hidden_offset: 0 → layer input, 1 → layer output. mem_norm_init: Initial value for mem_norm RMSNorm weight. uniform_num_selected: N for uniformly-spaced token selection. @@ -456,6 +503,8 @@ def load_metis_from_backbone( gated_delta_beta_init: Initial sigmoid value for gated-delta beta. qk_kernel_type: Feature map for kernelized q/k memory variants. metis_reweight_gamma: Gate weight for reweight blocks. + memory_on_sliding_layers: Whether sliding-attention layers carry + Metis memory. Gemma full-attention-only checkpoints set this false. Returns: (model, tokenizer) @@ -472,6 +521,7 @@ def load_metis_from_backbone( "metis_hyper_memory_type": metis_hyper_memory_type, "metis_local_memory_type": metis_local_memory_type, "update_ratio": update_ratio, + "forget_ratio": forget_ratio, "commit_hidden_offset": commit_hidden_offset, "mem_norm_init": mem_norm_init, "uniform_num_selected": uniform_num_selected, @@ -486,6 +536,7 @@ def load_metis_from_backbone( "gated_delta_beta_init": gated_delta_beta_init, "qk_kernel_type": qk_kernel_type, "metis_reweight_gamma": metis_reweight_gamma, + "memory_on_sliding_layers": memory_on_sliding_layers, }, ) @@ -497,10 +548,21 @@ def load_metis_from_backbone( model = MetisForCausalLM(config) print(f"[weight_utils] Loading backbone weights from {backbone_path} …") - backbone = AutoModelForCausalLM.from_pretrained(backbone_path, dtype=dtype) + if backbone_type_camel == "Gemma4": + backbone = _load_gemma4_pretrained(backbone_path, config, dtype) + else: + backbone = AutoModelForCausalLM.from_pretrained(backbone_path, dtype=dtype) mb = model.model.metis_backbone - _copy_module_state(mb.model, backbone.model) - _copy_module_state(mb.lm_head, backbone.lm_head) + if backbone_type_camel == "Gemma4": + text_model = _gemma4_text_model(backbone) + output_embeddings = backbone.get_output_embeddings() + if output_embeddings is None: + output_embeddings = text_model.embed_tokens + _copy_module_state(mb.model, text_model) + _copy_module_state(mb.lm_head, output_embeddings) + else: + _copy_module_state(mb.model, backbone.model) + _copy_module_state(mb.lm_head, backbone.lm_head) _init_learned_query_from_backbone(model) _init_hyper_memory_from_backbone(model)