Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
30 commits
Select commit Hold shift + click to select a range
ac6c680
Stage ①: Validate checkpoint format on entry (TODO-2, TODO-5, TODO-7)
amarquic Sep 9, 2026
f8902cf
Stage ②: Content-addressable cache key + checkpoint format detection …
amarquic Sep 9, 2026
df2f3ba
Stage ③: Eliminate Phase 1 shard scan, pass weight_map explicitly (TO…
amarquic Sep 9, 2026
3f4b72d
Stage ④: True pipeline — layout transform + DtypeConversion always ru…
amarquic Sep 10, 2026
4cd155c
Stage ⑤: Unified FusedExpertSplitCheckpointTransform with two-map app…
amarquic Sep 10, 2026
6aef8c8
TODO-3: Explicit ONNX→checkpoint key mapping via resolve_onnx_key()
amarquic Sep 10, 2026
32f4754
TODO-3: Explicit ONNX→checkpoint key mapping via resolve_onnx_key()
amarquic Sep 10, 2026
41913cb
Fix review findings: remove duplicate classes, fix get_consumed_keys()
amarquic Sep 10, 2026
f001ab6
Fix lm_head resolution, GptOss expert_parallel, unit test loader
amarquic Sep 15, 2026
966d11b
Remaining changes: test helpers, examples, MoE profiles, export utils
amarquic Sep 16, 2026
b0cd641
Add conftest.py changes
amarquic Sep 16, 2026
8ff69f5
minor typo fixed
amarquic Sep 16, 2026
31dc544
ruff format and minor changes
amarquic Sep 16, 2026
273aed1
Review fixes: ID map, hash inputs, flavours revert
amarquic Sep 18, 2026
bd7ebe9
Format weight-free test helpers
amarquic Sep 22, 2026
3664913
Implement staged checkpoint transform pipeline
amarquic Sep 24, 2026
ccaebc0
Document MLA retained-cache handling
amarquic Sep 24, 2026
10152b9
few necessary changes to run disagg with CB and remove guards for dis…
amarquic Sep 25, 2026
2765bb2
Format cache and MoE helpers
amarquic Sep 25, 2026
f7f223c
test(weight-free): cover disaggregated serving
amarquic Sep 25, 2026
64dfcc5
lint format update
amarquic Sep 25, 2026
3c0f04c
added example for disagg in weight-free and remove comment
amarquic Sep 25, 2026
43441e8
test(weight-free): validate disaggregated MDP compilation
amarquic Sep 27, 2026
6264ab0
fixed lint and format
amarquic Sep 28, 2026
d5da97d
updated example script and removed the comment in modelling auto
amarquic Sep 28, 2026
9bad57b
Added tests with hf vs qaic
quic-amitraj Sep 28, 2026
fb930ae
Addressed comments-1
quic-amitraj Sep 29, 2026
b6b722f
Addressed comments related to tests
quic-amitraj Sep 29, 2026
8aa9b7e
Fix retained-state transform test import
quic-amitraj Sep 29, 2026
0bd8c0c
fix: Dissag parity tests
quic-amitraj Oct 1, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
753 changes: 672 additions & 81 deletions QEfficient/base/checkpoint_transforms.py

Large diffs are not rendered by default.

94 changes: 62 additions & 32 deletions QEfficient/base/onnx_transforms.py
Original file line number Diff line number Diff line change
Expand Up @@ -303,14 +303,54 @@ class PreserveNestedCacheRetainedStateTransform(BaseOnnxTransform):
}
)

@staticmethod
def _scatter_sort_key(n) -> int:
output_name = n.output[0] if n.output else ""
if "key" in output_name:
return 0
if "value" in output_name:
return 1
return 2
@classmethod
def _resolve_kv_scatter_outputs(cls, fn, node, layer_idx: str) -> dict[str, str] | None:
"""Resolve nested KV scatter outputs from the function-call data flow.

ONNX function-call inputs bind positionally. Each cache scatter consumes
the cache buffer as its first input, so the matching scatter is identified
from that buffer's function argument rather than output names or node order.
"""
function_cache_inputs: dict[str, str] = {}
for input_index, input_name in enumerate(node.input):
match = cls._KV_INPUT_RE.match(input_name)
if match is None:
continue
kind, input_layer_idx = match.groups()
if input_layer_idx != layer_idx:
continue
if input_index >= len(fn.input):
raise ValueError(
f"Nested function '{fn.name}' has no input at position {input_index} for "
f"call-node cache input '{input_name}'."
)
function_cache_inputs[kind] = fn.input[input_index]

if set(function_cache_inputs) != {"key", "value"}:
return None

scatter_outputs = {}
for kind, function_input in function_cache_inputs.items():
writers = [
fn_node
for fn_node in fn.node
if fn_node.op_type in cls._SCATTER_OP_TYPES
and fn_node.input
and fn_node.input[0] == function_input
and fn_node.output
]
if len(writers) != 1:
writer_names = [
f"{writer.op_type}({writer.output[0] if writer.output else '<no-output>'})" for writer in writers
]
raise ValueError(
f"Could not uniquely resolve the nested past_{kind}.{layer_idx} cache writer in function "
f"'{fn.name}': expected one CtxScatter* with data input '{function_input}', "
f"found {len(writers)} ({writer_names})."
)
scatter_outputs[kind] = writers[0].output[0]

return scatter_outputs

@classmethod
def apply(cls, model: ModelProto) -> bool:
Expand All @@ -331,27 +371,6 @@ def apply(cls, model: ModelProto) -> bool:
if fn is None:
continue

# Collect scatter nodes that write back the KV cache.
# Sort by first-input name so the key scatter reliably precedes the
# value scatter: dynamo names function-body args generically (arg7_1
# etc.), so we sort by the scatter output name instead — dynamo
# preserves "key"/"value" in output tensor names even when input
# argument names are opaque.
scatter_nodes = [
fn_node for fn_node in fn.node if fn_node.op_type in cls._SCATTER_OP_TYPES and fn_node.output
]
if len(scatter_nodes) != 2:
Comment thread
vbaddi marked this conversation as resolved.
logger.debug(
"PreserveNestedCacheRetainedStateTransform: function '%s' has %d scatter node(s), expected 2 — skipping.",
node.op_type,
len(scatter_nodes),
)
continue
Comment thread
quic-amitraj marked this conversation as resolved.

scatter_nodes.sort(key=cls._scatter_sort_key)
# Only the first two scatter outputs map to key / value respectively.
scatter_outputs = [n.output[0] for n in scatter_nodes[:2]]

# Identify layer index from KV inputs on this call node.
layer_idx = None
kv_inputs = {}
Expand Down Expand Up @@ -379,10 +398,21 @@ def apply(cls, model: ModelProto) -> bool:
if not any(name in dangling_retained_outputs for name in desired_outputs):
continue

if not all(name in dangling_retained_outputs for name in desired_outputs):
raise ValueError(
f"Nested function '{fn.name}' has partially dangling KV retained-state outputs for layer "
f"{layer_idx}: expected both {desired_outputs}."
)

scatter_outputs = cls._resolve_kv_scatter_outputs(fn, node, layer_idx)
if scatter_outputs is None:
continue

# Expose scatter outputs in the function's output list, rename KV
# inputs and append retained-state output names to the call node —
# all in one pass over the two key/value pairs.
for kind, scatter_output, desired_output in zip(("key", "value"), scatter_outputs, desired_outputs):
# inputs and append retained-state output names to the call node.
# Both writers are resolved before this block, so graph rewiring is atomic.
for kind, desired_output in zip(("key", "value"), desired_outputs):
scatter_output = scatter_outputs[kind]
if scatter_output not in fn.output:
fn.output.append(scatter_output)
changed = True
Expand Down
18 changes: 18 additions & 0 deletions QEfficient/compile/mdp_generator.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@

import bisect
import logging
import re
from enum import Enum
from pathlib import Path
from typing import Any, Dict, List, Optional, Set, Tuple
Expand Down Expand Up @@ -286,6 +287,19 @@ def _get_layer_num(node_name: str) -> Optional[int]:
return None


def _get_layer_num_from_inputs(node_inputs: List[str]) -> Optional[int]:
"""Return a layer index from model parameter inputs on generic callsites.

Dynamo names repeated decoder-layer callsites generically, but their
weight inputs retain paths such as ``model.layers.3.self_attn.q_proj``.
"""
for input_name in node_inputs:
match = re.search(r"(?:^|\.)layers\.(\d+)(?:\.|$)", input_name)
if match:
return int(match.group(1))
return None


def _layer_partition_bounds(num_layers: int, num_partitions: int) -> List[int]:
"""Compute exclusive-upper-bound layer bounds for balanced pipeline partitioning.

Expand Down Expand Up @@ -499,6 +513,10 @@ def generate_disagg_mdp_partition_config(
continue

layer_num = _get_layer_num(node.name)
if layer_num is None:
# GPT-OSS/Dynamo repeated-subgraph callsites have generic names,
# while their model-layer parameter inputs retain ``layers.N``.
layer_num = _get_layer_num_from_inputs(list(node.input))
if layer_num is not None:
seen_first_layer = True
partition_idx = bisect.bisect_right(partition_bounds, layer_num)
Expand Down
88 changes: 70 additions & 18 deletions QEfficient/exporter/weight_free/checkpoint_key_resolver.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
#
# ----------------------------------------------------------------------------

import json
Comment thread
vbaddi marked this conversation as resolved.
from pathlib import Path
from typing import Dict, List, Optional

Expand Down Expand Up @@ -75,9 +76,18 @@ def _moe_weight_aliases(name: str) -> List[str]:
return aliases


def _router_gate_aliases(name: str) -> List[str]:
"""Return the legacy router/gate spelling for sparse-MoE router weights."""
if name.endswith(".mlp.gate.weight"):
return [name[: -len(".mlp.gate.weight")] + ".mlp.router.weight"]
if name.endswith(".mlp.router.weight"):
return [name[: -len(".mlp.router.weight")] + ".mlp.gate.weight"]
return []


def _find_checkpoint_key(candidates: List[str], checkpoint_index: Dict[str, str], onnx_name: str) -> Optional[str]:
"""Return the unique matching checkpoint key, or fail on ambiguous matches."""
seen = set()
seen: set = set()
matches = []
for candidate in candidates:
if candidate in seen:
Expand Down Expand Up @@ -108,14 +118,17 @@ def find_checkpoint_key(
onnx_name: str,
checkpoint_index: Dict[str, str],
backbone: nn.Module,
active_transform=None,
) -> Optional[str]:
"""Resolve an ONNX initializer name to its safetensors checkpoint key.

Most weights match directly. The fallback rules cover wrapper prefixes,
task-head/base-model checkpoint differences, and known HF/QEff MoE naming
differences without putting those details in the export orchestration path.
TODO(wf): Make this explicit model/layout mapping we should always know what key to expect in what case.
Resolution order:
1. Universal HF prefix rules (base_model., base_model_prefix).
2. Legacy sparse-MoE router/gate spelling fallback.
3. Transform-specific explicit mapping via resolve_onnx_key().
4. Legacy MoE weight aliases fallback for old checkpoints.
"""
# 1. Universal HF prefix rules
candidates = [onnx_name]
stripped = onnx_name.removeprefix("base_model.")
candidates.append(stripped)
Expand All @@ -127,16 +140,29 @@ def find_checkpoint_key(
if prefix and stripped.startswith(f"{prefix}."):
candidates.append(stripped[len(f"{prefix}.") :])

if ".mlp." in stripped:
candidates.append(stripped.replace(".mlp.", ".block_sparse_moe."))

if stripped.endswith(".mlp.gate.weight"):
candidates.append(stripped[: -len(".gate.weight")] + ".router.weight")

if stripped.endswith(".mlp.router.weight"):
candidates.append(stripped[: -len(".router.weight")] + ".gate.weight")

return _find_checkpoint_key(candidates, checkpoint_index, onnx_name)
key = _find_checkpoint_key(candidates, checkpoint_index, onnx_name)
if key is not None:
return key

# 2. Keep the historic sparse-MoE router/gate compatibility after exact lookup.
router_gate_candidates = [alias for candidate in candidates for alias in _router_gate_aliases(candidate)]
key = _find_checkpoint_key(router_gate_candidates, checkpoint_index, onnx_name)
if key is not None:
return key

# 3. Transform-specific explicit mapping
if active_transform is not None and hasattr(active_transform, "resolve_onnx_key"):
for candidate in [onnx_name, *_router_gate_aliases(onnx_name)]:
key = active_transform.resolve_onnx_key(candidate, checkpoint_index)
if key is not None:
return key

# 4. Legacy MoE weight aliases (kept for old prepared checkpoints)
return _find_checkpoint_key(
[alias for c in candidates for alias in _moe_weight_aliases(c)],
checkpoint_index,
onnx_name,
)


def promote_initializers_and_build_spec(onnx_program, model_ref: str, model_name: str, qeff_model) -> WeightSpec:
Expand All @@ -159,8 +185,8 @@ def promote_initializers_and_build_spec(onnx_program, model_ref: str, model_name
Specification mapping promoted ONNX inputs to checkpoint tensor locations.
"""
model_ir = onnx_program.model
parameter_names = {name for name, _ in qeff_model.model.named_parameters()}
buffer_names = {name for name, _ in qeff_model.model.named_buffers()}
parameter_names = {name for name, _ in qeff_model.model.named_parameters(remove_duplicate=False)}
buffer_names = {name for name, _ in qeff_model.model.named_buffers(remove_duplicate=False)}
model_names = parameter_names | buffer_names
tied_weight_map = {entry.alias: entry.canonical for entry in _collect_tied_weights(qeff_model.model)}
# named_parameters()/named_buffers() dedup tied tensors by identity, so a tied alias
Expand All @@ -179,14 +205,40 @@ def promote_initializers_and_build_spec(onnx_program, model_ref: str, model_name
for checkpoint_file in checkpoint_files
]
backbone = qeff_model.model.base_model if isinstance(qeff_model.model, PooledModel) else qeff_model.model

# Identify the active layout transform from the prepared checkpoint manifest.
# The manifest stores the active layout transform ID during centralized finalization.
# Reading from the manifest avoids re-running detection on the prepared checkpoint
# (which would fail — the prepared checkpoint has canonical output keys like
# moe_weights.gate, not the original per-expert keys that trigger detection).

from QEfficient.base.checkpoint_transforms import ( # noqa: PLC0415
CHECKPOINT_PREPARED_MANIFEST,
_find_transform_by_id,
)

active_transform = None
manifest_path = Path(model_ref) / CHECKPOINT_PREPARED_MANIFEST
if manifest_path.exists():
try:
manifest = json.loads(manifest_path.read_text())
transform_id = manifest.get("active_group", "none")
if transform_id and transform_id != "none":
active_transform = _find_transform_by_id(
transform_id,
getattr(qeff_model, "_checkpoint_transforms", []),
)
except (OSError, json.JSONDecodeError):
pass # no manifest → active_transform stays None, fallback to legacy aliases

promoted_inputs: List[WeightSpecInput] = []

for name, init_value in list(model_ir.graph.initializers.items()):
if name not in model_names:
continue

onnx_name = tied_weight_map.get(name, name)
checkpoint_key = find_checkpoint_key(onnx_name, checkpoint_index, backbone)
checkpoint_key = find_checkpoint_key(onnx_name, checkpoint_index, backbone, active_transform)
if checkpoint_key is None:
if _is_computed_initializer(onnx_name):
continue
Expand Down
Loading
Loading