Files
AHA-L2A-Qwen3-1.7B-repro/recipe/modeling_aha_qwen3.py
ModelHub XC c13d63b439 初始化项目,由ModelHub XC社区提供模型
Model: jiamingshan/AHA-L2A-Qwen3-1.7B-repro
Source: Original Platform
2026-07-21 11:06:13 +08:00

1227 lines
57 KiB
Python

"""
AHA-Qwen3: All-or-Here Attention applied to Qwen3.
Minimal modification to standard Qwen3 attention: adds a binary gate that
dynamically toggles between full attention and local attention. The persisted
``aha_router_granularity`` config selects either one gate per (token, KV-head)
or one gate per token shared by every head in the layer.
Based on: transformers.models.qwen3.modeling_qwen3 (Qwen3Attention)
Changes:
1. q_proj outputs either one extra gate value (``token``) or
``num_key_value_heads`` values (``token_kv_head``). The shared-token
value is broadcast to all KV/query heads before attention mixing.
2. Gate: sigmoid → hard threshold → STE for training
3. Double-mix: global attention + local window attention, blended by
gate (broadcast across the GQA group)
4. Auxiliary loss: λ * mean(gate_soft) to encourage sparsity
"""
import atexit
import json
import os
from dataclasses import dataclass
from pathlib import Path
from typing import Callable, Optional, Union
import torch
import torch.nn as nn
from safetensors.torch import load_file
from transformers import Qwen3Config, Qwen3ForCausalLM
from transformers.modeling_layers import GradientCheckpointingLayer
from transformers.models.qwen3.modeling_qwen3 import (
Qwen3Attention,
Qwen3DecoderLayer,
Qwen3Model,
Qwen3RMSNorm,
apply_rotary_pos_emb,
)
from transformers.cache_utils import Cache, DynamicCache
from transformers.masking_utils import create_causal_mask, create_sliding_window_causal_mask
from transformers.modeling_utils import ALL_ATTENTION_FUNCTIONS
from transformers.modeling_outputs import BaseModelOutputWithPast, CausalLMOutputWithPast
from transformers.processing_utils import Unpack
from transformers.utils.generic import TransformersKwargs
# AHA constants
AHA_WINDOW_SIZE = 1024
AHA_LAMBDA = 3e-4
AHA_DISTILL_WEIGHT = 0.0
AHA_CE_WEIGHT = 1.0
# Gate-density ceiling (hinge target). The aux loss only fires when
# ``mean(gate_soft) > AHA_GATE_TARGET``, so sparsity cannot drop
# below ``1 - AHA_GATE_TARGET`` in steady state regardless of how
# large ``aha_lambda`` is. Setting this to 1.0 recovers the legacy
# behaviour (unconditional downward pressure).
AHA_GATE_TARGET = 1.0
# DuoAttention baseline knobs. When ``aha_mode == "duo"`` the per-(token,
# kv_head) gate is replaced by a per-(layer, kv_head) learnable scalar
# (``duo_alpha``, 224 scalars total for Qwen3-0.6B), and the "local"
# branch becomes a streaming mask with ``duo_sink_size`` attention sinks
# plus a recent window of ``duo_recent_size``. This is a strict reduction
# of the dynamic AHA gate and reproduces the DuoAttention paper's head
# classifier (Xiao et al. 2024) on the Qwen3 backbone.
AHA_MODE = "dynamic"
DUO_SINK_SIZE = 64
DUO_RECENT_SIZE = 256
DUO_ALPHA_INIT = 1.0
# Local-branch mask kind for the dual-branch attention (global vs local).
# - "sliding_window": last ``aha_window_size`` tokens, no attention sinks.
# Historical default for ``aha_mode="dynamic"`` (Bv3 was trained with this).
# - "sink_recent": first ``duo_sink_size`` tokens + last ``duo_recent_size``
# tokens, matching the DuoAttention paper. Historical hardcoded default for
# ``aha_mode="duo"`` (still forced for duo regardless of this field).
# § 9.4.13.3 found that hot-starting a dynamic-mode ckpt from a duo ckpt with
# sliding_window local loses ~20pp LongBench retention vs the duo source,
# almost entirely from passage_retrieval_en (sinks are critical for long
# passage retrieval). Setting this to "sink_recent" for hot-started dynamic
# ckpts closes that architectural gap.
AHA_LOCAL_KIND = "sliding_window"
AHA_ROUTER_GRANULARITY = "token_kv_head"
def aha_router_output_size(config: Qwen3Config) -> int:
"""Return the number of native dynamic-router logits per token."""
granularity = getattr(config, "aha_router_granularity", AHA_ROUTER_GRANULARITY)
if granularity == "token":
return 1
if granularity == "token_kv_head":
return int(config.num_key_value_heads)
raise ValueError(f"Unsupported aha_router_granularity={granularity!r}")
class _AHAInferenceSparsityTracker:
"""Accumulate exact inference hard-route counts without changing outputs."""
def __init__(self) -> None:
# Keep counters as device scalars until process exit. Calling .item()
# once per layer and decode token serialized CUDA and made the
# diagnostic tracker materially slow generation.
self.sparse_decisions: dict[int, torch.Tensor] = {}
self.total_decisions: dict[int, int] = {}
self.phase_sparse_decisions: dict[str, torch.Tensor] = {}
self.phase_total_decisions: dict[str, int] = {}
raw_thresholds = os.environ.get("AHA_COUNTERFACTUAL_THRESHOLDS", "")
self.counterfactual_thresholds = tuple(
sorted({float(value) for value in raw_thresholds.split(",") if value.strip()})
)
if any(value < 0.0 or value > 1.0 for value in self.counterfactual_thresholds):
raise ValueError("AHA_COUNTERFACTUAL_THRESHOLDS values must be in [0, 1]")
self.counterfactual_sparse: dict[float, torch.Tensor] = {}
self.router_granularity: Optional[str] = None
self.native_router_decisions = 0
atexit.register(self.write_stats)
def update(
self,
gate_hard: torch.Tensor,
gate_soft: torch.Tensor,
layer_idx: int,
phase: str,
router_granularity: str,
native_router_width: int,
) -> None:
if self.router_granularity is None:
self.router_granularity = router_granularity
elif self.router_granularity != router_granularity:
raise RuntimeError(
"AHA sparsity tracker received mixed router granularities: "
f"{self.router_granularity!r} and {router_granularity!r}"
)
sparse = torch.count_nonzero(gate_hard == 0).to(torch.int64)
total = int(gate_hard.numel())
effective_width = int(gate_hard.shape[-1])
self.native_router_decisions += total // effective_width * native_router_width
if layer_idx not in self.sparse_decisions:
self.sparse_decisions[layer_idx] = sparse
else:
self.sparse_decisions[layer_idx].add_(sparse)
self.total_decisions[layer_idx] = self.total_decisions.get(layer_idx, 0) + total
if phase not in self.phase_sparse_decisions:
self.phase_sparse_decisions[phase] = sparse.clone()
else:
self.phase_sparse_decisions[phase].add_(sparse)
self.phase_total_decisions[phase] = self.phase_total_decisions.get(phase, 0) + total
if self.counterfactual_thresholds:
flat = gate_soft.detach().reshape(-1, 1)
thresholds = flat.new_tensor(self.counterfactual_thresholds).reshape(1, -1)
counts = (flat <= thresholds).sum(dim=0, dtype=torch.int64)
for threshold, count in zip(self.counterfactual_thresholds, counts.unbind()):
if threshold not in self.counterfactual_sparse:
self.counterfactual_sparse[threshold] = count
else:
self.counterfactual_sparse[threshold].add_(count)
def write_stats(self) -> None:
output = os.environ.get("AHA_SPARSITY_STATS_PATH")
if not output:
return
sparse_by_layer = {
layer_idx: int(value.item())
for layer_idx, value in self.sparse_decisions.items()
}
sparse = sum(sparse_by_layer.values())
total = sum(self.total_decisions.values())
per_layer = {}
for layer_idx in sorted(self.total_decisions):
layer_sparse = sparse_by_layer.get(layer_idx, 0)
layer_total = self.total_decisions[layer_idx]
per_layer[str(layer_idx)] = {
"sparse_decisions": layer_sparse,
"total_decisions": layer_total,
"sparsity": layer_sparse / layer_total if layer_total else None,
}
payload = {
"definition": (
"hard AHA gate zeros / all token x KV-head x layer decisions, "
"token-weighted over prefill and decode"
),
"router_granularity": self.router_granularity,
"native_router_decisions": self.native_router_decisions,
"effective_router_decisions": total,
"native_router_note": (
"Native decisions count learned gate logits before broadcast; effective "
"decisions always use token x KV-head x layer for matched sparsity."
),
"sparse_decisions": sparse,
"total_decisions": total,
"sparsity": sparse / total if total else None,
"full_attention_usage": 1.0 - sparse / total if total else None,
"per_layer": per_layer,
"by_phase": {
phase: {
"sparse_decisions": int(self.phase_sparse_decisions[phase].item()),
"total_decisions": phase_total,
"sparsity": int(self.phase_sparse_decisions[phase].item()) / phase_total
if phase_total else None,
"full_attention_usage": 1.0
- int(self.phase_sparse_decisions[phase].item()) / phase_total
if phase_total else None,
}
for phase, phase_total in sorted(self.phase_total_decisions.items())
},
}
if self.counterfactual_thresholds:
payload["counterfactual_sparsity_by_threshold"] = {
f"{threshold:.6g}": {
"sparse_decisions": int(self.counterfactual_sparse[threshold].item()),
"total_decisions": total,
"sparsity": int(self.counterfactual_sparse[threshold].item()) / total
if total else None,
}
for threshold in self.counterfactual_thresholds
}
payload["counterfactual_note"] = (
"Routing counts only; benchmark quality must be measured in a separate "
"run with AHA_GATE_HARD_THRESHOLD set to the selected threshold."
)
path = Path(output)
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text(json.dumps(payload, indent=2) + "\n")
_AHA_INFERENCE_SPARSITY_TRACKER: Optional[_AHAInferenceSparsityTracker] = None
def _track_aha_inference_sparsity(
gate_hard: torch.Tensor,
gate_soft: torch.Tensor,
layer_idx: int,
phase: str,
router_granularity: str,
native_router_width: int,
) -> None:
global _AHA_INFERENCE_SPARSITY_TRACKER
if not os.environ.get("AHA_SPARSITY_STATS_PATH"):
return
if _AHA_INFERENCE_SPARSITY_TRACKER is None:
_AHA_INFERENCE_SPARSITY_TRACKER = _AHAInferenceSparsityTracker()
_AHA_INFERENCE_SPARSITY_TRACKER.update(
gate_hard.detach(), gate_soft.detach(), layer_idx, phase,
router_granularity, native_router_width,
)
def _get_optional_nonnegative_int_env(name: str) -> Optional[int]:
raw = os.environ.get(name, "").strip()
if not raw:
return None
value = int(raw)
if value < 0:
raise ValueError(f"{name} must be non-negative, got {value}")
return value
def _apply_low_alpha_force_full_heads_from_env(model: nn.Module) -> None:
"""Auto-protect the weakest static-full heads in duo_dynamic mode.
``AHA_FORCE_LOW_ALPHA_FULL_HEADS=K`` selects the K smallest
``full_attention_heads`` values among heads that static Duo still kept
full (alpha > 0.5). These are the fragile boundary full heads: raising the
dynamic threshold can close them, but static Duo's alpha says they should
not be fully discarded. The selected heads are OR-ed with any manually
supplied ``AHA_FORCE_FULL_HEADS`` mask.
"""
k = _get_optional_nonnegative_int_env("AHA_FORCE_LOW_ALPHA_FULL_HEADS")
if not k:
return
if getattr(model.config, "aha_mode", "dynamic") != "duo_dynamic":
return
candidates: list[tuple[float, int, int]] = []
with torch.no_grad():
for layer_idx, layer in enumerate(model.model.layers):
attn = layer.self_attn
alpha = getattr(attn, "full_attention_heads", None)
if alpha is None:
continue
for head_idx, value in enumerate(alpha.detach().float().cpu().tolist()):
if value > 0.5:
candidates.append((float(value), layer_idx, head_idx))
selected = sorted(candidates)[: min(k, len(candidates))]
if not selected:
return
selected_by_layer: dict[int, list[int]] = {}
for _, layer_idx, head_idx in selected:
selected_by_layer.setdefault(layer_idx, []).append(head_idx)
with torch.no_grad():
for layer_idx, heads in selected_by_layer.items():
mask = model.model.layers[layer_idx].self_attn._aha_force_full_heads_mask
for head_idx in heads:
mask[head_idx] = True
model.config.aha_force_low_alpha_full_heads = int(k)
formatted = ",".join(f"{layer}:{head}" for _, layer, head in selected)
print(
f"[AHA] AHA_FORCE_LOW_ALPHA_FULL_HEADS={k} selected {len(selected)} "
f"duo_dynamic heads: {formatted}",
flush=True,
)
class AHAQwen3Config(Qwen3Config):
"""Qwen3Config with AHA-specific parameters."""
model_type = "aha_qwen3"
def __init__(
self,
aha_window_size=AHA_WINDOW_SIZE,
aha_lambda=AHA_LAMBDA,
aha_distill_weight=AHA_DISTILL_WEIGHT,
aha_ce_weight=AHA_CE_WEIGHT,
aha_gate_target=AHA_GATE_TARGET,
aha_reg_weight=-1.0,
aha_force_gate_value=None,
aha_mode=AHA_MODE,
duo_sink_size=DUO_SINK_SIZE,
duo_recent_size=DUO_RECENT_SIZE,
duo_alpha_init=DUO_ALPHA_INIT,
aha_local_kind=AHA_LOCAL_KIND,
aha_router_granularity=AHA_ROUTER_GRANULARITY,
**kwargs,
):
super().__init__(**kwargs)
self.aha_window_size = aha_window_size
self.aha_lambda = aha_lambda
# Per-layer attention distillation weight. When > 0 and model is
# training, each layer contributes ``mean(((1-g)·(global-local))**2)``
# to the total loss. This is a dense per-(token, head) signal that
# pushes the gate toward 1 where SWA is insufficient to approximate
# full attention. Cost: ~0 extra forward compute (global/local are
# already computed for the mix).
self.aha_distill_weight = aha_distill_weight
# Language-modeling CE weight for the training objective. Default 1.0
# keeps the standard SFT behaviour. Set to 0.0 to drop CE entirely
# and train the gate purely from ``aux + distill`` (useful when the
# backbone is frozen and we only want to shape the gate). Evaluation
# always reports unweighted ``ce_loss`` regardless of this setting.
self.aha_ce_weight = aha_ce_weight
# Direct Duo-style sparsity regularizer. When >= 0, this replaces the
# legacy hinge term and uses ``aha_reg_weight * mean(gate_soft)``.
# Default -1 keeps older checkpoints on their original hinge objective.
self.aha_reg_weight = float(aha_reg_weight)
# Optional override used by Duo-style dynamic distillation. When set
# to 1.0, dynamic mode behaves as a full-attention teacher for one
# forward; None keeps the learned gate.
self.aha_force_gate_value = aha_force_gate_value
# Target ceiling on mean(gate_soft). Aux loss is hinge-shaped:
# ``λ · max(0, mean(gate_soft) - τ)``. Setting τ < 1.0 guarantees
# the gate cannot drop below ``1 - τ`` local-attention fraction
# in steady state -- the failure mode we observed with the legacy
# unconditional aux, where ``λ = 1.0`` drove mean(gate_hard) to
# 0.005 (99.5% local) and killed long-context accuracy
# (gsm8k 74% -> 19%). Default 1.0 is a no-op for backward
# compatibility with checkpoints trained on the old loss.
self.aha_gate_target = aha_gate_target
# DuoAttention-style baseline. ``dynamic`` keeps the original per-
# (token, kv_head) MLP gate. ``duo`` replaces it with a per-(layer,
# kv_head) learnable scalar and swaps the local branch's mask for a
# streaming (sink + recent) pattern. ``duo_dynamic`` keeps the Duo
# static mask for KV-capacity planning, but applies a dynamic AHA gate
# only on the static-full heads: effective_full = duo_full AND dyn_gate.
# Sink / recent sizes are fixed at train time; set ``duo_sink_size=0``
# for the sink-ablation run.
if aha_mode not in ("dynamic", "duo", "duo_dynamic"):
raise ValueError(
f"aha_mode must be 'dynamic', 'duo', or 'duo_dynamic', got {aha_mode!r}"
)
self.aha_mode = aha_mode
if aha_router_granularity not in ("token", "token_kv_head"):
raise ValueError(
"aha_router_granularity must be 'token' or 'token_kv_head', "
f"got {aha_router_granularity!r}"
)
# Persist this even for static Duo checkpoints. Dynamic AHA uses it to
# choose the native q_proj gate-row count; old configs omit the field
# and therefore retain the historical token_kv_head behavior.
self.aha_router_granularity = aha_router_granularity
self.duo_sink_size = int(duo_sink_size)
self.duo_recent_size = int(duo_recent_size)
self.duo_alpha_init = float(duo_alpha_init)
if aha_local_kind not in ("sliding_window", "sink_recent"):
raise ValueError(
f"aha_local_kind must be 'sliding_window' or 'sink_recent', "
f"got {aha_local_kind!r}"
)
# When aha_mode="duo" the local branch is always sink+recent regardless
# of this field (forced to the historical DuoAttention behaviour). The
# field only controls the dynamic mode's local mask.
self.aha_local_kind = aha_local_kind
def _build_streaming_causal_mask(
seq_len_q: int,
seq_len_kv: int,
sink_size: int,
recent_size: int,
device: torch.device,
dtype: torch.dtype,
kv_offset: int = 0,
) -> torch.Tensor:
"""Construct a streaming attention mask (sink + recent, causal).
Returns an additive mask of shape ``[1, 1, Q, K]`` where
disallowed positions are ``finfo(dtype).min`` and allowed positions
are ``0``. A query at position ``q`` (absolute in the sequence,
with ``kv_offset`` tokens already in the KV cache) attends to a key
at position ``k`` iff ``k <= q_abs`` AND (``k < sink_size`` OR
``q_abs - k < recent_size``).
"""
q_abs = torch.arange(seq_len_q, device=device).unsqueeze(-1) + kv_offset # [Q, 1]
k_abs = torch.arange(seq_len_kv, device=device).unsqueeze(0) # [1, K]
causal = k_abs <= q_abs
sink = k_abs < sink_size if sink_size > 0 else torch.zeros_like(causal)
recent = (q_abs - k_abs) < recent_size
allowed = causal & (sink | recent)
mask = torch.zeros(seq_len_q, seq_len_kv, dtype=dtype, device=device)
mask.masked_fill_(~allowed, torch.finfo(dtype).min)
return mask.unsqueeze(0).unsqueeze(0)
class AHAQwen3Attention(nn.Module):
"""Qwen3 attention with AHA gate for per-head local/global routing."""
def __init__(self, config: AHAQwen3Config, layer_idx: int):
super().__init__()
self.config = config
self.layer_idx = layer_idx
self.head_dim = getattr(config, "head_dim", config.hidden_size // config.num_attention_heads)
self.num_key_value_groups = config.num_attention_heads // config.num_key_value_heads
self.scaling = self.head_dim ** -0.5
self.attention_dropout = config.attention_dropout
self.is_causal = True
self.aha_window_size = config.aha_window_size
# q_proj: extra router outputs for either one shared per-token gate or
# one per-(token, KV-head) gate.
# DUO mode uses a per-(layer, kv_head) static scalar instead, so
# q_proj keeps its vanilla Qwen3 shape and is binary-compatible
# with the base model checkpoint.
self._has_dyn_gate_logits = getattr(config, "aha_mode", "dynamic") in (
"dynamic",
"duo_dynamic",
)
self.aha_router_outputs = aha_router_output_size(config)
q_out_dim = config.num_attention_heads * self.head_dim
if self._has_dyn_gate_logits:
q_out_dim += self.aha_router_outputs
self.q_proj = nn.Linear(
config.hidden_size,
q_out_dim,
bias=config.attention_bias,
)
self.k_proj = nn.Linear(
config.hidden_size, config.num_key_value_heads * self.head_dim, bias=config.attention_bias
)
self.v_proj = nn.Linear(
config.hidden_size, config.num_key_value_heads * self.head_dim, bias=config.attention_bias
)
self.o_proj = nn.Linear(
config.num_attention_heads * self.head_dim, config.hidden_size, bias=config.attention_bias
)
# Qwen3 applies q_norm/k_norm per head_dim (not full projection)
self.q_norm = Qwen3RMSNorm(self.head_dim, eps=config.rms_norm_eps)
self.k_norm = Qwen3RMSNorm(self.head_dim, eps=config.rms_norm_eps)
# DuoAttention baseline: per-(layer, kv_head) learnable scalar.
# Populated in ``duo`` and ``duo_dynamic``. In ``duo_dynamic`` this
# scalar is the frozen Duo capacity mask; the dynamic gate is only
# allowed to close heads where this scalar says "full". Registered
# as ``full_attention_heads`` to match the name used in the
# official DuoAttention repo (``duo_attn.patch.llama``), so
# per-head statistics (e.g. retrieval heads vs streaming heads)
# can be read back with the same tooling.
if getattr(config, "aha_mode", "dynamic") in ("duo", "duo_dynamic"):
init = float(getattr(config, "duo_alpha_init", 1.0))
self.full_attention_heads = nn.Parameter(
torch.full((config.num_key_value_heads,), init, dtype=torch.float32)
)
else:
self.full_attention_heads = None
force_full_mask = torch.zeros(config.num_key_value_heads, dtype=torch.bool)
for raw in os.environ.get("AHA_FORCE_FULL_HEADS", "").split(","):
raw = raw.strip()
if not raw:
continue
layer_s, head_s = raw.split(":", 1)
if int(layer_s) == int(layer_idx):
force_full_mask[int(head_s)] = True
self.register_buffer(
"_aha_force_full_heads_mask",
force_full_mask,
persistent=False,
)
def forward(
self,
hidden_states: torch.Tensor,
position_embeddings: tuple[torch.Tensor, torch.Tensor],
attention_mask: Optional[dict], # dict with "global" and "local" keys
past_key_values: Optional[Cache] = None,
cache_position: Optional[torch.LongTensor] = None,
**kwargs,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, Optional[torch.Tensor]]:
input_shape = hidden_states.shape[:-1]
hidden_shape = (*input_shape, -1, self.head_dim)
# === Q projection + (optional) gate extraction ===
q_proj_out = self.q_proj(hidden_states)
if self._has_dyn_gate_logits:
query_states, gate = torch.split(
q_proj_out,
[
self.config.num_attention_heads * self.head_dim,
self.aha_router_outputs,
],
dim=-1,
)
else:
query_states = q_proj_out
# DUO mode: use a zero placeholder so downstream dtype/device
# checks stay unchanged. The actual gate comes from the
# per-kv-head scalar ``full_attention_heads`` below.
gate = torch.zeros(
*input_shape, self.config.num_key_value_heads,
dtype=query_states.dtype, device=query_states.device,
)
# Qwen3-style: norm then reshape then transpose
query_states = self.q_norm(query_states.view(hidden_shape)).transpose(1, 2)
key_states = self.k_norm(self.k_proj(hidden_states).view(hidden_shape)).transpose(1, 2)
value_states = self.v_proj(hidden_states).view(hidden_shape).transpose(1, 2)
cos, sin = position_embeddings
query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin)
if past_key_values is not None:
cache_kwargs = {"sin": sin, "cos": cos, "cache_position": cache_position}
key_states, value_states = past_key_values.update(
key_states, value_states, self.layer_idx, cache_kwargs
)
# === Gate computation ===
# gate shape: (B, Q, H_kv) — per-(token, kv_head) in dynamic mode,
# or per-kv_head broadcast to (B, Q, H_kv) in DUO mode.
# At inference, threshold is read from env var AHA_GATE_HARD_THRESHOLD
# so users can trade sparsity for quality without re-training.
# Default 0.5 keeps training-time behaviour. Lowering to e.g. 0.35
# lets more (token, kv_head) pairs route through global attention,
# which recovers long-CoT quality at the cost of sparsity.
aha_mode = getattr(self.config, "aha_mode", "dynamic")
if aha_mode == "duo":
# DuoAttention: alpha is per-(layer, kv_head) learnable scalar.
# gate_soft = clamp(alpha, 0, 1); gate_hard = (alpha > 0.5).
# Both broadcast to (B, Q, H_kv) for a uniform downstream API
# (so the distill / aux / blend paths do not need to branch).
qhead_alpha = getattr(self, "_aha_qhead_full_attention_heads", None)
if qhead_alpha is not None:
# Eval-only oracle diagnostic: route at Q-head granularity
# instead of KV-head/GQA-group granularity. This does not
# represent a KV-capacity-saving deployment point; it isolates
# whether GQA grouping itself is the quality bottleneck.
alpha = qhead_alpha.to(gate.dtype).clamp(0.0, 1.0)
gate_soft = alpha.view(1, 1, -1).expand(*input_shape, -1)
import os as _os
_thr = float(_os.environ.get("AHA_GATE_HARD_THRESHOLD", "0.5"))
gate_hard = (alpha > _thr).to(gate.dtype).view(1, 1, -1).expand(*input_shape, -1)
gate_ste = gate_hard
else:
alpha = self.full_attention_heads.to(gate.dtype).clamp(0.0, 1.0)
gate_soft = alpha.view(1, 1, -1).expand(*input_shape, -1)
if self.training:
gate_hard = (alpha > 0.5).to(gate.dtype).view(1, 1, -1).expand(*input_shape, -1)
# Straight-through estimator through the clamped alpha so
# L1 / distill gradients flow into the scalar.
gate_ste = gate_soft
else:
import os as _os
_thr = float(_os.environ.get("AHA_GATE_HARD_THRESHOLD", "0.5"))
gate_hard = (alpha > _thr).to(gate.dtype).view(1, 1, -1).expand(*input_shape, -1)
gate_ste = gate_hard
elif aha_mode == "duo_dynamic":
# Duo + AHA-on-full-heads:
# duo_full_mask = 1 for static Duo full/retrieval heads
# dyn_gate = per-token AHA decision inside those heads
# effective_full = duo_full_mask AND dyn_gate
#
# Duo streaming heads are locked to local/streaming forever. This
# preserves Duo's KV-capacity saving: only static-full heads need a
# full KV cache. ``aha_force_gate_value=1`` therefore reproduces
# the static Duo deployment point exactly, not all-full attention.
alpha = self.full_attention_heads.to(gate.dtype).clamp(0.0, 1.0)
duo_full = (alpha > 0.5).to(gate.dtype).view(1, 1, -1)
dyn_soft = torch.sigmoid(gate)
force_gate_value = getattr(self.config, "aha_force_gate_value", None)
if force_gate_value is not None:
dyn_soft = torch.full_like(dyn_soft, float(force_gate_value))
gate_soft = dyn_soft * duo_full
gate_hard = gate_soft
gate_ste = gate_soft
elif self.training:
train_thr = float(os.environ.get("AHA_TRAIN_GATE_HARD_THRESHOLD", "0.5"))
gate_soft = dyn_soft * duo_full
dyn_hard = (dyn_soft > train_thr).to(gate.dtype)
gate_hard = dyn_hard * duo_full
gate_ste = gate_hard + (gate_soft - gate_soft.detach())
else:
import os as _os
_thr = float(_os.environ.get("AHA_GATE_HARD_THRESHOLD", "0.5"))
gate_soft = dyn_soft * duo_full
dyn_hard = (dyn_soft > _thr).to(gate.dtype)
gate_hard = dyn_hard * duo_full
gate_ste = gate_hard
force_full = self._aha_force_full_heads_mask.to(gate.dtype).view(1, 1, -1)
if bool(force_full.any().item()):
force_full = force_full * duo_full
gate_hard = torch.maximum(gate_hard, force_full)
gate_ste = torch.maximum(gate_ste, force_full)
else:
gate_soft = torch.sigmoid(gate)
force_gate_value = getattr(self.config, "aha_force_gate_value", None)
if force_gate_value is not None:
force = float(force_gate_value)
gate_soft = torch.full_like(gate_soft, force)
gate_hard = torch.full_like(gate_soft, force)
gate_ste = gate_soft
elif self.training:
train_thr = float(os.environ.get("AHA_TRAIN_GATE_HARD_THRESHOLD", "0.5"))
gate_hard = (gate_soft > train_thr).to(gate.dtype)
gate_ste = gate_hard + (gate_soft - gate_soft.detach())
else:
import os as _os
_thr = float(_os.environ.get("AHA_GATE_HARD_THRESHOLD", "0.5"))
gate_hard = (gate_soft > _thr).to(gate.dtype)
gate_ste = gate_hard
# Normalize the downstream contract to effective
# token x KV-head decisions. In shared-token mode there is only one
# learned logit per token, but regularization, sparsity accounting, and
# attention mixing all see the same KV-head denominator as the
# head-granular arm.
if gate_soft.shape[-1] == 1 and self.config.num_key_value_heads != 1:
effective_shape = (*gate_soft.shape[:-1], self.config.num_key_value_heads)
gate_soft = gate_soft.expand(effective_shape)
gate_hard = gate_hard.expand(effective_shape)
gate_ste = gate_ste.expand(effective_shape)
if (
aha_mode in ("duo", "duo_dynamic")
and not self.training
and os.environ.get("AHA_DUO_PREFILL_FULL", "0") == "1"
and query_states.shape[-2] == key_states.shape[-2]
):
# Diagnostic parity with the official DuoAttention HF eval path:
# initial prefill is full attention; streaming heads are only
# localized during subsequent decode steps.
gate_hard = torch.ones_like(gate_hard)
gate_ste = torch.ones_like(gate_ste)
if not self.training:
phase = (
"decode"
if hidden_states.shape[1] == 1 and key_states.shape[-2] > 1
else "prefill"
)
_track_aha_inference_sparsity(
gate_hard, gate_soft, self.layer_idx, phase,
getattr(
self.config,
"aha_router_granularity",
AHA_ROUTER_GRANULARITY,
),
self.aha_router_outputs if self._has_dyn_gate_logits
else self.config.num_key_value_heads,
)
# === Attention: double_mix (global + local, blended by gate) ===
from transformers.models.qwen3.modeling_qwen3 import eager_attention_forward
attention_interface: Callable = eager_attention_forward
if self.config._attn_implementation != "eager":
attention_interface = ALL_ATTENTION_FUNCTIONS[self.config._attn_implementation]
# Global attention (full context)
global_attn_output, _ = attention_interface(
self, query_states, key_states, value_states,
attention_mask["global"],
dropout=0.0 if not self.training else self.attention_dropout,
scaling=self.scaling,
**kwargs,
)
if getattr(self.config, "_aha_teacher_full_fastpath", False):
attn_output = global_attn_output.reshape(*input_shape, -1).contiguous()
attn_output = self.o_proj(attn_output)
return attn_output, gate_soft, gate_hard, None
# Local / streaming branch. The local mask was pre-built by the
# model's ``forward`` either as sliding-window (legacy dynamic mode
# default) or sink+recent (duo mode, and dynamic mode when
# ``aha_local_kind="sink_recent"``). Only pass ``sliding_window`` to
# the SDPA backend when the mask is the sliding-window kind, otherwise
# the backend will double-mask and wipe out the sink positions.
local_attn_kwargs = dict(
dropout=0.0 if not self.training else self.attention_dropout,
scaling=self.scaling,
)
local_kind = "sink_recent" if aha_mode in ("duo", "duo_dynamic") else getattr(
self.config, "aha_local_kind", "sliding_window"
)
if local_kind == "sliding_window":
local_attn_kwargs["sliding_window"] = self.aha_window_size
local_attn_output, _ = attention_interface(
self, query_states, key_states, value_states,
attention_mask["local"],
**local_attn_kwargs,
**kwargs,
)
# Blend: normal gate_ste is (B, Q, H_kv); broadcast to (B, Q, H, 1)
# by repeating each kv_head decision. The optional qhead oracle above
# already emits (B, Q, H), so it bypasses this broadcast.
if gate_ste.shape[-1] == self.config.num_attention_heads:
g_bqh = gate_ste
else:
g_bqh = gate_ste.repeat_interleave(self.num_key_value_groups, dim=-1)
attn_output = (
global_attn_output * g_bqh.unsqueeze(-1)
+ local_attn_output * (1 - g_bqh.unsqueeze(-1))
)
# Per-layer attention distillation signal. Returns a scalar so the
# outer model can aggregate across layers. When gate → 1 the term
# vanishes; when gate → 0 on positions where global/local disagree
# strongly, the gradient pushes gate back toward 1.
#
# ``diff`` is ``detach``-ed so distill only trains the gate; it does
# not reshape attention outputs through the backbone. This also
# prevents the squared-difference graph from retaining activations
# for all 28 layers' global/local branches, which was the source of
# OOM when gradient checkpointing is disabled on the custom forward.
distill_per_layer: Optional[torch.Tensor] = None
if self.training and getattr(self.config, "aha_distill_weight", 0.0) > 0.0:
with torch.no_grad():
diff_sq = (global_attn_output - local_attn_output).float() ** 2
one_minus_g_sq = ((1.0 - g_bqh.float()) ** 2).unsqueeze(-1)
distill_per_layer = (one_minus_g_sq * diff_sq).mean()
attn_output = attn_output.reshape(*input_shape, -1).contiguous()
attn_output = self.o_proj(attn_output)
return attn_output, gate_soft, gate_hard, distill_per_layer
class AHAQwen3DecoderLayer(GradientCheckpointingLayer):
"""Decoder layer with gradient-checkpointing support.
Inheriting from ``GradientCheckpointingLayer`` lets HuggingFace Trainer's
``gradient_checkpointing_enable`` take effect on this layer without any
explicit ``_gradient_checkpointing_func`` call in the outer forward.
"""
def __init__(self, config: AHAQwen3Config, layer_idx: int):
super().__init__()
self.hidden_size = config.hidden_size
self.self_attn = AHAQwen3Attention(config=config, layer_idx=layer_idx)
from transformers.models.qwen3.modeling_qwen3 import Qwen3MLP
self.mlp = Qwen3MLP(config)
self.input_layernorm = Qwen3RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
self.post_attention_layernorm = Qwen3RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
def forward(
self,
hidden_states: torch.Tensor,
attention_mask: Optional[dict] = None,
position_ids: Optional[torch.LongTensor] = None,
past_key_values: Optional[Cache] = None,
use_cache: Optional[bool] = False,
cache_position: Optional[torch.LongTensor] = None,
position_embeddings: Optional[tuple[torch.Tensor, torch.Tensor]] = None,
**kwargs,
):
residual = hidden_states
hidden_states = self.input_layernorm(hidden_states)
hidden_states, gate_soft, gate_hard, distill_per_layer = self.self_attn(
hidden_states=hidden_states,
attention_mask=attention_mask,
position_ids=position_ids,
past_key_values=past_key_values,
use_cache=use_cache,
cache_position=cache_position,
position_embeddings=position_embeddings,
**kwargs,
)
hidden_states = residual + hidden_states
residual = hidden_states
hidden_states = self.post_attention_layernorm(hidden_states)
hidden_states = self.mlp(hidden_states)
hidden_states = residual + hidden_states
return hidden_states, gate_soft, gate_hard, distill_per_layer
@dataclass
class AHACausalLMOutputWithPast(CausalLMOutputWithPast):
"""Causal LM output with AHA loss breakdown and gate density stats (training / eval with labels)."""
ce_loss: Optional[torch.FloatTensor] = None
gate_aux_loss: Optional[torch.FloatTensor] = None
distill_loss: Optional[torch.FloatTensor] = None
gate_soft_mean: Optional[torch.FloatTensor] = None
gate_hard_mean: Optional[torch.FloatTensor] = None
class AHAModelOutputWithPast(BaseModelOutputWithPast):
def __init__(self, last_hidden_state, past_key_values=None, hidden_states=None, attentions=None,
all_gate_soft=None, all_gate_hard=None, all_distill_per_layer=None):
super().__init__(last_hidden_state=last_hidden_state, past_key_values=past_key_values,
hidden_states=hidden_states, attentions=attentions)
self.all_gate_soft = all_gate_soft
self.all_gate_hard = all_gate_hard
# Tuple of per-layer distill scalars, or empty tuple when distillation is off.
self.all_distill_per_layer = all_distill_per_layer if all_distill_per_layer is not None else ()
class AHAQwen3Model(Qwen3Model):
config_class = AHAQwen3Config
def __init__(self, config: AHAQwen3Config):
from transformers import PreTrainedModel
PreTrainedModel.__init__(self, config)
self.padding_idx = config.pad_token_id
self.vocab_size = config.vocab_size
self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx)
self.layers = nn.ModuleList(
[AHAQwen3DecoderLayer(config, layer_idx) for layer_idx in range(config.num_hidden_layers)]
)
self.norm = Qwen3RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
from transformers.models.qwen3.modeling_qwen3 import Qwen3RotaryEmbedding
self.rotary_emb = Qwen3RotaryEmbedding(config=config)
self.gradient_checkpointing = False
self.post_init()
def forward(
self,
input_ids: Optional[torch.LongTensor] = None,
attention_mask: Optional[torch.Tensor] = None,
position_ids: Optional[torch.LongTensor] = None,
past_key_values: Optional[Cache] = None,
inputs_embeds: Optional[torch.FloatTensor] = None,
use_cache: Optional[bool] = None,
cache_position: Optional[torch.LongTensor] = None,
**kwargs,
) -> AHAModelOutputWithPast:
if (input_ids is None) ^ (inputs_embeds is not None):
raise ValueError("You must specify exactly one of input_ids or inputs_embeds")
if inputs_embeds is None:
inputs_embeds = self.embed_tokens(input_ids)
if use_cache and past_key_values is None:
past_key_values = DynamicCache(config=self.config)
if cache_position is None:
past_seen_tokens = past_key_values.get_seq_length() if past_key_values is not None else 0
cache_position = torch.arange(
past_seen_tokens, past_seen_tokens + inputs_embeds.shape[1], device=inputs_embeds.device
)
if position_ids is None:
position_ids = cache_position.unsqueeze(0)
# Build two masks: global (full causal) and local (sliding window
# for dynamic mode, streaming sink+recent for DUO mode).
orig_sliding_window = getattr(self.config, "sliding_window", None)
self.config.sliding_window = self.config.aha_window_size
mask_kwargs = {
"config": self.config,
"input_embeds": inputs_embeds,
"attention_mask": attention_mask,
"cache_position": cache_position,
"past_key_values": past_key_values,
"position_ids": position_ids,
}
global_mask = create_causal_mask(**mask_kwargs)
aha_mode = getattr(self.config, "aha_mode", "dynamic")
# duo / duo_dynamic force sink_recent for backward compat. dynamic mode
# reads the local kind from config (default sliding_window for legacy
# Bv3 ckpts; hot-started ckpts from duo set sink_recent — see
# §9.4.13.3).
local_kind = "sink_recent" if aha_mode in ("duo", "duo_dynamic") else getattr(
self.config, "aha_local_kind", "sliding_window"
)
if local_kind == "sink_recent":
seq_len_q = inputs_embeds.shape[1]
past_seen = past_key_values.get_seq_length() if past_key_values is not None else 0
seq_len_kv = seq_len_q + past_seen
local_mask = _build_streaming_causal_mask(
seq_len_q=seq_len_q,
seq_len_kv=seq_len_kv,
sink_size=int(getattr(self.config, "duo_sink_size", 64)),
recent_size=int(getattr(self.config, "duo_recent_size", 256)),
device=inputs_embeds.device,
dtype=inputs_embeds.dtype,
kv_offset=past_seen,
)
else:
local_mask = create_sliding_window_causal_mask(**mask_kwargs)
causal_mask_dict = {"global": global_mask, "local": local_mask}
self.config.sliding_window = orig_sliding_window
hidden_states = inputs_embeds
position_embeddings = self.rotary_emb(hidden_states, position_ids)
all_gate_soft = ()
all_gate_hard = ()
all_distill_per_layer: tuple[torch.Tensor, ...] = ()
for decoder_layer in self.layers:
hidden_states, gate_soft, gate_hard, distill_per_layer = decoder_layer(
hidden_states,
attention_mask=causal_mask_dict,
position_ids=position_ids,
past_key_values=past_key_values,
use_cache=use_cache,
cache_position=cache_position,
position_embeddings=position_embeddings,
**kwargs,
)
all_gate_soft += (gate_soft,)
all_gate_hard += (gate_hard,)
if distill_per_layer is not None:
all_distill_per_layer += (distill_per_layer,)
hidden_states = self.norm(hidden_states)
return AHAModelOutputWithPast(
last_hidden_state=hidden_states,
past_key_values=past_key_values if use_cache else None,
all_gate_soft=all_gate_soft,
all_gate_hard=all_gate_hard,
all_distill_per_layer=all_distill_per_layer,
)
class AHAQwen3ForCausalLM(Qwen3ForCausalLM):
config_class = AHAQwen3Config
_tied_weights_keys = ["lm_head.weight"]
def __init__(self, config: AHAQwen3Config):
# Skip Qwen3ForCausalLM.__init__ which creates Qwen3Model
# Go to PreTrainedModel.__init__ directly
from transformers import PreTrainedModel
PreTrainedModel.__init__(self, config)
self.model = AHAQwen3Model(config)
self.vocab_size = config.vocab_size
self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
self.post_init()
def forward(
self,
input_ids: Optional[torch.LongTensor] = None,
attention_mask: Optional[torch.Tensor] = None,
position_ids: Optional[torch.LongTensor] = None,
past_key_values: Optional[Cache] = None,
inputs_embeds: Optional[torch.FloatTensor] = None,
labels: Optional[torch.LongTensor] = None,
use_cache: Optional[bool] = None,
cache_position: Optional[torch.LongTensor] = None,
logits_to_keep: Union[int, torch.Tensor] = 0,
**kwargs,
) -> AHACausalLMOutputWithPast:
outputs = self.model(
input_ids=input_ids,
attention_mask=attention_mask,
position_ids=position_ids,
past_key_values=past_key_values,
inputs_embeds=inputs_embeds,
use_cache=use_cache,
cache_position=cache_position,
**kwargs,
)
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.lm_head(hidden_states[:, slice_indices, :])
loss = None
ce_loss = None
gate_aux_loss = None
distill_loss = None
gate_soft_mean = None
gate_hard_mean = None
if labels is not None:
from torch.nn import CrossEntropyLoss
shift_logits = logits[..., :-1, :].contiguous()
shift_labels = labels[..., 1:].contiguous()
loss_fct = CrossEntropyLoss()
ce_loss = loss_fct(shift_logits.view(-1, self.vocab_size), shift_labels.view(-1))
all_gate_soft = outputs.all_gate_soft
all_gate_hard = outputs.all_gate_hard
if labels is not None and len(all_gate_soft) > 0 and len(all_gate_hard) > 0:
gate_soft_stack = torch.stack(
[g.to(hidden_states.device) for g in all_gate_soft], dim=-1
).float()
gate_hard_stack = torch.stack(
[g.to(hidden_states.device).float() for g in all_gate_hard], dim=-1
)
shift_gate_soft = gate_soft_stack[:, :-1, :].contiguous()
shift_gate_hard = gate_hard_stack[:, :-1, :].contiguous()
# Match the CE/PPL population exactly. The old unmasked mean
# counted padded sequence positions as real gate decisions, which
# made the training sparsity curve disagree with inference-time
# sparsity on variable-length batches.
valid_gate = shift_labels.ne(-100)
if attention_mask is not None:
valid_gate = valid_gate & attention_mask[:, 1:].to(torch.bool)
if bool(valid_gate.any()):
valid_gate = valid_gate[..., None, None].to(shift_gate_soft.dtype)
normalizer = (
valid_gate.sum()
* shift_gate_soft.shape[-2]
* shift_gate_soft.shape[-1]
)
gate_soft_mean = (shift_gate_soft * valid_gate).sum() / normalizer
gate_hard_mean = (shift_gate_hard * valid_gate).sum() / normalizer
else:
gate_soft_mean = shift_gate_soft.mean()
gate_hard_mean = shift_gate_hard.mean()
# Per-layer distillation loss: mean over layers of
# ``mean(((1-g)·(global-local))**2)``. Populated only when training
# with ``config.aha_distill_weight > 0`` (see AHAQwen3Attention).
all_distill = outputs.all_distill_per_layer
if self.training and len(all_distill) > 0:
distill_loss = torch.stack(list(all_distill)).mean()
# === Loss assembly ===
# Eval path: standard LM loss for compatibility with HF metrics.
# Train path: weighted sum of (ce, gate_aux, distill) so callers can
# zero any term by setting its weight to 0 (e.g. ``aha_ce_weight=0``
# to train the gate purely from ``aux + distill``).
if not self.training:
loss = ce_loss
else:
ce_weight = float(getattr(self.config, "aha_ce_weight", 1.0))
distill_weight = float(getattr(self.config, "aha_distill_weight", 0.0))
terms: list[torch.Tensor] = []
if ce_loss is not None and ce_weight > 0.0:
terms.append(ce_weight * ce_loss)
if gate_soft_mean is not None:
reg_weight = float(getattr(self.config, "aha_reg_weight", -1.0))
if reg_weight >= 0.0:
# Duo-style direct sparsity term:
# alpha/gate near 1 means full/global attention, so penalize
# mean gate usage directly.
gate_aux_loss = reg_weight * gate_soft_mean
else:
# Legacy hinge aux: penalise only the excess above the target
# ceiling. Below ``τ`` the gradient is zero, so CE and
# distill freely shape where the gate opens without an
# unbounded downward push from aux. See AHAQwen3Config
# for the rationale and the collapse that motivated it.
gate_target = float(getattr(self.config, "aha_gate_target", 1.0))
excess = torch.clamp(gate_soft_mean - gate_target, min=0.0)
gate_aux_loss = self.config.aha_lambda * excess
terms.append(gate_aux_loss)
if distill_loss is not None and distill_weight > 0.0:
terms.append(distill_weight * distill_loss)
if terms:
loss = terms[0]
for t in terms[1:]:
loss = loss + t
return AHACausalLMOutputWithPast(
loss=loss,
logits=logits,
past_key_values=outputs.past_key_values,
ce_loss=ce_loss,
gate_aux_loss=gate_aux_loss,
distill_loss=distill_loss,
gate_soft_mean=gate_soft_mean,
gate_hard_mean=gate_hard_mean,
)
@classmethod
def from_pretrained_qwen3(cls, qwen3_path: str, aha_window_size: int = AHA_WINDOW_SIZE,
aha_lambda: float = AHA_LAMBDA,
aha_distill_weight: float = AHA_DISTILL_WEIGHT,
aha_ce_weight: float = AHA_CE_WEIGHT,
aha_gate_target: float = AHA_GATE_TARGET,
aha_reg_weight: float = -1.0,
aha_force_gate_value=None,
aha_mode: str = AHA_MODE,
duo_sink_size: int = DUO_SINK_SIZE,
duo_recent_size: int = DUO_RECENT_SIZE,
duo_alpha_init: float = DUO_ALPHA_INIT,
aha_local_kind: str = AHA_LOCAL_KIND,
aha_router_granularity: str = AHA_ROUTER_GRANULARITY,
**kwargs):
"""Load a pretrained Qwen3 model and convert to AHA-Qwen3.
Initializes the gate weights in q_proj randomly (the extra num_heads outputs).
All other weights are copied from the pretrained model.
"""
# Load original Qwen3
original = Qwen3ForCausalLM.from_pretrained(qwen3_path, **kwargs)
original_config = original.config
# Create AHA config
config_dict = original_config.to_dict()
config_dict["aha_window_size"] = aha_window_size
config_dict["aha_lambda"] = aha_lambda
config_dict["aha_distill_weight"] = aha_distill_weight
config_dict["aha_ce_weight"] = aha_ce_weight
config_dict["aha_gate_target"] = aha_gate_target
config_dict["aha_reg_weight"] = aha_reg_weight
config_dict["aha_force_gate_value"] = aha_force_gate_value
config_dict["aha_mode"] = aha_mode
config_dict["duo_sink_size"] = duo_sink_size
config_dict["duo_recent_size"] = duo_recent_size
config_dict["duo_alpha_init"] = duo_alpha_init
config_dict["aha_local_kind"] = aha_local_kind
config_dict["aha_router_granularity"] = aha_router_granularity
config_dict["model_type"] = "aha_qwen3"
aha_config = AHAQwen3Config(**config_dict)
aha_config._attn_implementation = kwargs.get("attn_implementation", "sdpa")
# Create AHA model
aha_model = cls(aha_config)
# Copy weights
# embed_tokens, lm_head, norm
aha_model.model.embed_tokens.load_state_dict(original.model.embed_tokens.state_dict())
aha_model.lm_head.load_state_dict(original.lm_head.state_dict())
aha_model.model.norm.load_state_dict(original.model.norm.state_dict())
aha_model.model.rotary_emb.load_state_dict(original.model.rotary_emb.state_dict())
# Copy per-layer weights
num_heads = original_config.num_attention_heads
head_dim = aha_config.head_dim
for i in range(original_config.num_hidden_layers):
orig_layer = original.model.layers[i]
aha_layer = aha_model.model.layers[i]
# MLP + norms: direct copy
aha_layer.mlp.load_state_dict(orig_layer.mlp.state_dict())
aha_layer.input_layernorm.load_state_dict(orig_layer.input_layernorm.state_dict())
aha_layer.post_attention_layernorm.load_state_dict(orig_layer.post_attention_layernorm.state_dict())
# Attention: k_proj, v_proj, o_proj, q_norm, k_norm — direct copy
aha_attn = aha_layer.self_attn
orig_attn = orig_layer.self_attn
aha_attn.k_proj.load_state_dict(orig_attn.k_proj.state_dict())
aha_attn.v_proj.load_state_dict(orig_attn.v_proj.state_dict())
aha_attn.o_proj.load_state_dict(orig_attn.o_proj.state_dict())
aha_attn.q_norm.load_state_dict(orig_attn.q_norm.state_dict())
aha_attn.k_norm.load_state_dict(orig_attn.k_norm.state_dict())
# q_proj:
# - dynamic / duo_dynamic mode: copy original Q rows, init extra
# gate rows random (~0.5).
# - duo mode: q_proj is identical shape to base Qwen3 (no gate
# rows), so direct load_state_dict works.
orig_q_weight = orig_attn.q_proj.weight.data # (num_heads*head_dim, hidden_size)
aha_q_weight = aha_attn.q_proj.weight.data
if aha_attn._has_dyn_gate_logits:
aha_q_weight[:num_heads * head_dim, :] = orig_q_weight
nn.init.normal_(aha_q_weight[num_heads * head_dim:, :], mean=0.0, std=0.01)
if orig_attn.q_proj.bias is not None and aha_attn.q_proj.bias is not None:
aha_attn.q_proj.bias.data[:num_heads * head_dim] = orig_attn.q_proj.bias.data
nn.init.zeros_(aha_attn.q_proj.bias.data[num_heads * head_dim:])
else:
aha_attn.q_proj.load_state_dict(orig_attn.q_proj.state_dict())
del original
return aha_model
@classmethod
def from_pretrained_aha(cls, checkpoint_path: str, **kwargs):
"""Load an AHA-Qwen3 checkpoint saved by `trainer.save_model`.
Supports both single-file (``model.safetensors``) and multi-shard
(``model.safetensors.index.json`` + ``model-*-of-*.safetensors``)
layouts. The latter is automatically used by HF when a checkpoint
exceeds the single-file size threshold (e.g. Qwen3-1.7B+).
"""
config = AHAQwen3Config.from_pretrained(checkpoint_path)
config._attn_implementation = kwargs.get(
"attn_implementation",
getattr(config, "_attn_implementation", "sdpa"),
)
model = cls(config)
single_path = os.path.join(checkpoint_path, "model.safetensors")
index_path = os.path.join(checkpoint_path, "model.safetensors.index.json")
if os.path.exists(single_path):
state_dict = load_file(single_path)
elif os.path.exists(index_path):
with open(index_path) as f:
shard_index = json.load(f)
shard_files = sorted(set(shard_index["weight_map"].values()))
state_dict = {}
for shard in shard_files:
state_dict.update(load_file(os.path.join(checkpoint_path, shard)))
else:
raise FileNotFoundError(
f"No safetensors file or index found in {checkpoint_path}. "
f"Expected either model.safetensors or model.safetensors.index.json."
)
missing_keys, unexpected_keys = model.load_state_dict(state_dict, strict=False)
allowed_missing = {"lm_head.weight"} if getattr(config, "tie_word_embeddings", False) else set()
extra_missing = set(missing_keys) - allowed_missing
if extra_missing or unexpected_keys:
raise RuntimeError(
f"Unexpected checkpoint mismatch. missing={sorted(extra_missing)}, "
f"unexpected={sorted(unexpected_keys)}"
)
if "lm_head.weight" in missing_keys:
model.tie_weights()
_apply_low_alpha_force_full_heads_from_env(model)
torch_dtype = kwargs.get("torch_dtype")
if torch_dtype is not None:
model = model.to(dtype=torch_dtype)
return model