diff --git a/QEfficient/base/onnx_transforms.py b/QEfficient/base/onnx_transforms.py
index 05f8773025..74c6ed0894 100644
--- a/QEfficient/base/onnx_transforms.py
+++ b/QEfficient/base/onnx_transforms.py
@@ -5,6 +5,9 @@
#
# ----------------------------------------------------------------------------
+import copy
+import hashlib
+import json
import logging
import os
import re
@@ -14,7 +17,7 @@
import numpy as np
import onnx
import torch
-from onnx import ModelProto, TensorProto, external_data_helper, numpy_helper
+from onnx import AttributeProto, ModelProto, TensorProto, external_data_helper, numpy_helper
from QEfficient.customop.ctx_scatter_gather import (
CtxChunkScatterBatch,
@@ -421,6 +424,59 @@ def _rename_graph_inputs_bulk(graph: onnx.GraphProto, rename_map: Dict[str, str]
return changed
+class CanonicalizeWhileLoopInitialConditionTransform(BaseOnnxTransform):
+ """Make Dynamo while-loop initial conditions ONNX constants.
+
+ ``torch.while_loop`` lowers its initial condition through the captured
+ condition graph. ONNX permits that value to be dynamic, but QAIC requires
+ the initial ``Loop`` condition to be constant. A scalar ``Constant`` node
+ supplies that initial ``true`` value; the loop body still carries and
+ updates the dynamic condition for subsequent iterations.
+ """
+
+ _COND_GRAPH_MARKER = "while_loop_cond_graph"
+
+ @classmethod
+ def _constant_node(cls, name: str) -> onnx.NodeProto:
+ value = numpy_helper.from_array(np.asarray(True, dtype=np.bool_), name=name)
+ return onnx.helper.make_node("Constant", [], [name], value=value)
+
+ @classmethod
+ def _rewrite_nodes(cls, nodes) -> bool:
+ changed = False
+ index = 0
+ while index < len(nodes):
+ node = nodes[index]
+ if node.op_type == "Loop" and len(node.input) > 1 and node.input[1]:
+ condition_name = node.input[1]
+ condition_producer = next(
+ (candidate for candidate in nodes[:index] if condition_name in candidate.output), None
+ )
+ if condition_producer is not None and cls._COND_GRAPH_MARKER in condition_producer.op_type:
+ constant_name = f"{condition_name}_constant"
+ nodes.insert(index, cls._constant_node(constant_name))
+ node.input[1] = constant_name
+ changed = True
+ index += 1
+
+ for attribute in node.attribute:
+ if attribute.HasField("g") and cls._rewrite_nodes(attribute.g.node):
+ changed = True
+ for graph in attribute.graphs:
+ if cls._rewrite_nodes(graph.node):
+ changed = True
+ index += 1
+ return changed
+
+ @classmethod
+ def apply(cls, model: ModelProto, **kwargs) -> bool:
+ del kwargs
+ changed = cls._rewrite_nodes(model.graph.node)
+ for function in model.functions:
+ if cls._rewrite_nodes(function.node):
+ changed = True
+ return changed
+
class RenameRepeatedSubgraphTransform(BaseOnnxTransform):
"""Rename dynamo repeated_subgraph function names to model-specific layer class names.
@@ -447,6 +503,8 @@ def _iter_all_nodes(cls, nodes):
for attr in node.attribute:
if attr.HasField("g"):
yield from cls._iter_all_nodes(attr.g.node)
+ for graph in attr.graphs:
+ yield from cls._iter_all_nodes(graph.node)
@staticmethod
def _rename_op_types(nodes, old_to_new: Dict[str, str]) -> None:
@@ -516,6 +574,293 @@ def apply(cls, model: ModelProto, target_classnames: Optional[List[str]] = None,
return True
+class DeduplicateRepeatedSubgraphTransform(BaseOnnxTransform):
+ """Collapse structurally identical repeated decoder local functions.
+
+ Dynamo subfunction export can emit one FunctionProto per decoder layer, even
+ when those function bodies differ only by local SSA names or layer-indexed
+ formal parameter names. ONNX local-function calls bind inputs and outputs by
+ position, so duplicate call nodes can safely target the first structurally
+ equivalent function without changing graph-level tensor names.
+ """
+
+ _NUMERIC_SUFFIX_RE = re.compile(r"^(?P.+)_(?P\d+)$")
+
+ @classmethod
+ def apply(cls, model: ModelProto, target_classnames: Optional[List[str]] = None, **kwargs) -> bool:
+ target_classnames = [name for name in (target_classnames or []) if name]
+ candidates = cls._collect_candidate_functions(model, target_classnames)
+ if len(candidates) < 2:
+ return False
+
+ input_dim_signatures = cls._function_input_dim_signatures(model, candidates)
+ fingerprint_to_canonical = {}
+ duplicate_to_canonical = {}
+ functions_to_remove = set()
+
+ for _, fn in candidates:
+ fingerprint = cls._function_fingerprint(fn, input_dim_signatures.get(fn.name))
+ canonical = fingerprint_to_canonical.get(fingerprint)
+ if canonical is None:
+ fingerprint_to_canonical[fingerprint] = fn
+ continue
+ duplicate_to_canonical[fn.name] = canonical.name
+ functions_to_remove.add(fn.name)
+
+ if not duplicate_to_canonical:
+ return False
+
+ cls._rewrite_function_calls(model.graph.node, duplicate_to_canonical)
+ for fn in model.functions:
+ cls._rewrite_function_calls(fn.node, duplicate_to_canonical)
+
+ kept_functions = [fn for fn in model.functions if fn.name not in functions_to_remove]
+ del model.functions[:]
+ model.functions.extend(kept_functions)
+ return True
+
+ @classmethod
+ def _collect_candidate_functions(cls, model: ModelProto, target_classnames: List[str]):
+ called_function_names = cls._called_function_names(model)
+ candidates = []
+ for fn in model.functions:
+ if fn.name not in called_function_names:
+ continue
+ order = cls._candidate_order(fn.name, target_classnames)
+ if order is None:
+ continue
+ candidates.append((order, fn))
+ candidates.sort(key=lambda item: item[0])
+ return candidates
+
+ @classmethod
+ def _candidate_order(cls, name: str, target_classnames: List[str]) -> Optional[Tuple[int, int, str]]:
+ for pattern_index, pattern in enumerate(RenameRepeatedSubgraphTransform._REPEATED_SUBGRAPH_PATTERNS):
+ match = pattern.match(name)
+ if match:
+ return pattern_index, int(match.group(1)), name
+
+ for class_index, class_name in enumerate(target_classnames):
+ if name == class_name:
+ return 100 + class_index, 0, name
+ match = cls._NUMERIC_SUFFIX_RE.match(name)
+ if match and match.group("base") == class_name:
+ return 100 + class_index, int(match.group("idx")), name
+
+ return None
+
+ @classmethod
+ def _called_function_names(cls, model: ModelProto) -> set[str]:
+ function_names = {fn.name for fn in model.functions}
+ called = set()
+ for node in cls._iter_all_nodes(model.graph.node):
+ if node.op_type in function_names:
+ called.add(node.op_type)
+ for fn in model.functions:
+ for node in cls._iter_all_nodes(fn.node):
+ if node.op_type in function_names:
+ called.add(node.op_type)
+ return called
+
+ @staticmethod
+ def _rewrite_function_calls(nodes, old_to_new) -> None:
+ for node in DeduplicateRepeatedSubgraphTransform._iter_all_nodes(nodes):
+ new_op_type = old_to_new.get(node.op_type)
+ if new_op_type is None:
+ continue
+ node.op_type = new_op_type
+
+ @staticmethod
+ def _iter_all_nodes(nodes):
+ yield from RenameRepeatedSubgraphTransform._iter_all_nodes(nodes)
+
+ @classmethod
+ def _function_fingerprint(
+ cls, fn: onnx.FunctionProto, input_dim_signatures: Optional[Dict[int, tuple]] = None
+ ) -> str:
+ state = cls._new_value_state(fn.input)
+ payload = {
+ "domain": fn.domain,
+ "inputs": [state["value"](name) for name in fn.input],
+ "outputs": [state["value"](name) for name in fn.output],
+ "opsets": sorted((opset.domain, opset.version) for opset in fn.opset_import),
+ "nodes": [cls._node_key(node, state, input_dim_signatures) for node in fn.node],
+ }
+ encoded = json.dumps(payload, sort_keys=True, separators=(",", ":")).encode("utf-8")
+ return hashlib.sha256(encoded).hexdigest()
+
+ @staticmethod
+ def _new_value_state(inputs):
+ input_index = {name: idx for idx, name in enumerate(inputs)}
+ value_map = {name: f"arg{idx}" for idx, name in enumerate(inputs)}
+ counter = {"value": 0}
+
+ def value(name: str) -> str:
+ if not name:
+ return ""
+ if name not in value_map:
+ value_map[name] = f"tmp{counter['value']}"
+ counter["value"] += 1
+ return value_map[name]
+
+ def assign(name: str, canonical_value: str) -> None:
+ if name:
+ value_map[name] = canonical_value
+
+ return {"value": value, "assign": assign, "input_index": input_index}
+
+ @classmethod
+ def _node_key(cls, node: onnx.NodeProto, state, input_dim_signatures: Optional[Dict[int, tuple]] = None) -> tuple:
+ shape_dim_key = cls._shape_dim_key(node, state, input_dim_signatures)
+ if shape_dim_key is not None:
+ canonical_output = cls._shape_dim_value_name(shape_dim_key)
+ state["assign"](node.output[0], canonical_output)
+ return node.domain, "ShapeDim", shape_dim_key
+
+ squeezed_shape_dim = cls._squeezed_shape_dim_value(node, state)
+ if squeezed_shape_dim is not None:
+ state["assign"](node.output[0], squeezed_shape_dim)
+ return node.domain, "SqueezeShapeDim", squeezed_shape_dim
+
+ return (
+ node.domain,
+ node.op_type,
+ tuple(state["value"](name) for name in node.input),
+ tuple(state["value"](name) for name in node.output),
+ tuple(cls._attribute_key(attr, state) for attr in sorted(node.attribute, key=lambda attr: attr.name)),
+ )
+
+ @classmethod
+ def _shape_dim_key(
+ cls, node: onnx.NodeProto, state, input_dim_signatures: Optional[Dict[int, tuple]]
+ ) -> Optional[tuple]:
+ if node.op_type != "Shape" or len(node.input) != 1 or len(node.output) != 1 or input_dim_signatures is None:
+ return None
+
+ input_index = state["input_index"].get(node.input[0])
+ if input_index is None:
+ return None
+
+ dims = input_dim_signatures.get(input_index)
+ if not dims:
+ return None
+
+ attrs = {attr.name: attr.i for attr in node.attribute if attr.name in {"start", "end"}}
+ start = attrs.get("start", 0)
+ end = attrs.get("end", len(dims))
+ if start < 0:
+ start += len(dims)
+ if end < 0:
+ end += len(dims)
+ if start < 0 or end > len(dims) or end - start != 1:
+ return None
+
+ dim = dims[start]
+ if dim is None:
+ return None
+ return "shape_dim", dim
+
+ @staticmethod
+ def _shape_dim_value_name(shape_dim_key: tuple) -> str:
+ _, dim = shape_dim_key
+ return "shape_dim_" + "_".join(str(part) for part in dim)
+
+ @classmethod
+ def _squeezed_shape_dim_value(cls, node: onnx.NodeProto, state) -> Optional[str]:
+ if node.op_type != "Squeeze" or len(node.input) != 1 or len(node.output) != 1:
+ return None
+ input_value = state["value"](node.input[0])
+ if not input_value.startswith("shape_dim_"):
+ return None
+ return f"squeezed_{input_value}"
+
+ @classmethod
+ def _attribute_key(cls, attr: onnx.AttributeProto, state) -> tuple:
+ attr_type = attr.type
+ if attr_type == AttributeProto.FLOAT:
+ return attr.name, "f", attr.f
+ if attr_type == AttributeProto.INT:
+ return attr.name, "i", attr.i
+ if attr_type == AttributeProto.STRING:
+ return attr.name, "s", attr.s.decode("utf-8", errors="replace")
+ if attr_type == AttributeProto.FLOATS:
+ return attr.name, "floats", tuple(attr.floats)
+ if attr_type == AttributeProto.INTS:
+ return attr.name, "ints", tuple(attr.ints)
+ if attr_type == AttributeProto.STRINGS:
+ return attr.name, "strings", tuple(s.decode("utf-8", errors="replace") for s in attr.strings)
+ if attr_type == AttributeProto.TENSOR:
+ return attr.name, "t", cls._tensor_key(attr.t)
+ if attr_type == AttributeProto.TENSORS:
+ return attr.name, "tensors", tuple(cls._tensor_key(tensor) for tensor in attr.tensors)
+ if attr_type == AttributeProto.GRAPH:
+ return attr.name, "g", cls._graph_key(attr.g)
+ if attr_type == AttributeProto.GRAPHS:
+ return attr.name, "graphs", tuple(cls._graph_key(graph) for graph in attr.graphs)
+ return attr.name, "raw", attr.SerializeToString().hex()
+
+ @classmethod
+ def _graph_key(cls, graph: onnx.GraphProto) -> tuple:
+ inputs = [value.name for value in graph.input]
+ state = cls._new_value_state(inputs)
+ return (
+ tuple(state["value"](value.name) for value in graph.input),
+ tuple(state["value"](value.name) for value in graph.output),
+ tuple(cls._tensor_key(tensor) for tensor in graph.initializer),
+ tuple(cls._node_key(node, state) for node in graph.node),
+ )
+
+ @classmethod
+ def _function_input_dim_signatures(cls, model: ModelProto, candidates) -> Dict[str, Dict[int, tuple]]:
+ function_by_name = {fn.name: fn for _, fn in candidates}
+ value_shapes = cls._graph_value_shape_signatures(model.graph)
+ observed_shapes: Dict[str, Dict[int, set]] = {name: {} for name in function_by_name}
+
+ for node in cls._iter_all_nodes(model.graph.node):
+ fn = function_by_name.get(node.op_type)
+ if fn is None:
+ continue
+ for idx, input_name in enumerate(node.input[: len(fn.input)]):
+ shape = value_shapes.get(input_name)
+ if shape is None:
+ continue
+ observed_shapes[node.op_type].setdefault(idx, set()).add(shape)
+
+ input_dim_signatures = {}
+ for function_name, shapes_by_input in observed_shapes.items():
+ stable_shapes = {idx: next(iter(shapes)) for idx, shapes in shapes_by_input.items() if len(shapes) == 1}
+ input_dim_signatures[function_name] = stable_shapes
+ return input_dim_signatures
+
+ @staticmethod
+ def _graph_value_shape_signatures(graph: onnx.GraphProto) -> Dict[str, tuple]:
+ value_shapes = {}
+
+ def dim_key(dim, index):
+ if dim.dim_param:
+ return "sym", dim.dim_param
+ if dim.HasField("dim_value"):
+ return "value", dim.dim_value
+ return "unknown", index
+
+ for value_info in list(graph.input) + list(graph.value_info) + list(graph.output):
+ tensor_type = value_info.type.tensor_type
+ if not tensor_type.HasField("shape"):
+ continue
+ value_shapes[value_info.name] = tuple(dim_key(dim, idx) for idx, dim in enumerate(tensor_type.shape.dim))
+
+ for initializer in graph.initializer:
+ value_shapes[initializer.name] = tuple(("value", dim) for dim in initializer.dims)
+
+ return value_shapes
+
+ @staticmethod
+ def _tensor_key(tensor: onnx.TensorProto) -> tuple:
+ tensor = copy.deepcopy(tensor)
+ tensor.name = ""
+ return tensor.data_type, tuple(tensor.dims), tensor.SerializeToString().hex()
+
+
class AdapterWeightsToInputsTransform(BaseOnnxTransform):
@classmethod
@@ -615,6 +960,133 @@ def apply(cls, model: ModelProto) -> bool:
return transformed
+class LocalizeFunctionReduceSumAxesTransform(BaseOnnxTransform):
+ """Move constant ReduceSum axes from function arguments into function bodies."""
+
+ _INTEGER_TENSOR_TYPES = {
+ TensorProto.INT8,
+ TensorProto.INT16,
+ TensorProto.INT32,
+ TensorProto.INT64,
+ TensorProto.UINT8,
+ TensorProto.UINT16,
+ TensorProto.UINT32,
+ TensorProto.UINT64,
+ }
+
+ @classmethod
+ def apply(cls, model: ModelProto) -> bool:
+ transformed = False
+ graph_constants = cls._collect_graph_constants(model.graph)
+
+ for function in model.functions:
+ function_inputs = list(function.input)
+ if not function_inputs or cls._has_nested_call_site(model, function):
+ continue
+
+ call_sites = cls._find_graph_call_sites(model, function)
+ if not call_sites:
+ continue
+
+ axes_inputs = cls._find_reduce_sum_axes_inputs(function, function_inputs)
+ for axes_name, reduce_nodes in sorted(
+ axes_inputs.items(), key=lambda item: function_inputs.index(item[0]), reverse=True
+ ):
+ formal_index = function_inputs.index(axes_name)
+ axes_tensor = cls._resolve_shared_axes_tensor(call_sites, graph_constants, formal_index)
+ if axes_tensor is None:
+ continue
+
+ local_axes_name = cls._insert_axes_constant(function, axes_name, axes_tensor)
+ for reduce_node in reduce_nodes:
+ reduce_node.input[1] = local_axes_name
+
+ del function.input[formal_index]
+ for call_node in call_sites:
+ del call_node.input[formal_index]
+ transformed = True
+
+ return transformed
+
+ @classmethod
+ def _find_reduce_sum_axes_inputs(cls, function, function_inputs):
+ formal_inputs = set(function_inputs)
+ axes_inputs = {}
+ unsafe_inputs = set()
+
+ for node in function.node:
+ for input_index, input_name in enumerate(node.input):
+ if input_name not in formal_inputs:
+ continue
+ if node.op_type == "ReduceSum" and input_index == 1:
+ axes_inputs.setdefault(input_name, []).append(node)
+ else:
+ unsafe_inputs.add(input_name)
+
+ for input_name in unsafe_inputs:
+ axes_inputs.pop(input_name, None)
+ return axes_inputs
+
+ @staticmethod
+ def _find_graph_call_sites(model, function):
+ return [node for node in model.graph.node if node.op_type == function.name and node.domain == function.domain]
+
+ @staticmethod
+ def _has_nested_call_site(model, function):
+ for caller_function in model.functions:
+ for node in caller_function.node:
+ if node.op_type == function.name and node.domain == function.domain:
+ return True
+ return False
+
+ @classmethod
+ def _collect_graph_constants(cls, graph):
+ constants = {}
+ for initializer in graph.initializer:
+ if initializer.data_type in cls._INTEGER_TENSOR_TYPES:
+ constants[initializer.name] = tuple(numpy_helper.to_array(initializer).reshape(-1).tolist())
+ for node in graph.node:
+ if node.op_type != "Constant" or not node.output:
+ continue
+ for attribute in node.attribute:
+ if attribute.name == "value" and attribute.type == AttributeProto.TENSOR:
+ tensor = attribute.t
+ if tensor.data_type in cls._INTEGER_TENSOR_TYPES:
+ constants[node.output[0]] = tuple(numpy_helper.to_array(tensor).reshape(-1).tolist())
+ return constants
+
+ @classmethod
+ def _resolve_shared_axes_tensor(cls, call_sites, graph_constants, formal_index):
+ values = []
+ for call_node in call_sites:
+ if formal_index >= len(call_node.input):
+ return None
+ value = graph_constants.get(call_node.input[formal_index])
+ if value is None:
+ return None
+ values.append(value)
+ if not values or any(value != values[0] for value in values[1:]):
+ return None
+ return values[0]
+
+ @staticmethod
+ def _insert_axes_constant(function, formal_name, values):
+ base_name = f"{formal_name}_localized"
+ used_names = set(function.input) | set(function.output)
+ for node in function.node:
+ used_names.update(node.input)
+ used_names.update(node.output)
+ local_name = base_name
+ suffix = 0
+ while local_name in used_names:
+ suffix += 1
+ local_name = f"{base_name}_{suffix}"
+
+ tensor = numpy_helper.from_array(np.asarray(values, dtype=np.int64), name=local_name)
+ function.node.insert(0, onnx.helper.make_node("Constant", [], [local_name], value=tensor))
+ return local_name
+
+
class OnnxTransformPipeline(BaseOnnxTransform):
"""Pipeline to apply multiple ONNX transformations in sequence."""
@@ -691,9 +1163,20 @@ def _set_external_data(tensor, file_name):
if RenameWsubNodesTransform in requested:
applied[RenameWsubNodesTransform] = RenameWsubNodesTransform.apply(model)
+ if LocalizeFunctionReduceSumAxesTransform in requested:
+ applied[LocalizeFunctionReduceSumAxesTransform] = LocalizeFunctionReduceSumAxesTransform.apply(model)
+
if PreserveNestedCacheRetainedStateTransform in requested:
applied[PreserveNestedCacheRetainedStateTransform] = PreserveNestedCacheRetainedStateTransform.apply(model)
+ if CanonicalizeWhileLoopInitialConditionTransform in requested:
+ applied[CanonicalizeWhileLoopInitialConditionTransform] = (
+ CanonicalizeWhileLoopInitialConditionTransform.apply(model)
+ )
+
+ if DeduplicateRepeatedSubgraphTransform in requested:
+ applied[DeduplicateRepeatedSubgraphTransform] = DeduplicateRepeatedSubgraphTransform.apply(model, **kwargs)
+
if RenameRepeatedSubgraphTransform in requested:
applied[RenameRepeatedSubgraphTransform] = RenameRepeatedSubgraphTransform.apply(model, **kwargs)
diff --git a/QEfficient/exporter/weight_free/checkpoint_key_resolver.py b/QEfficient/exporter/weight_free/checkpoint_key_resolver.py
index 62d255a8f8..2345230bfb 100644
--- a/QEfficient/exporter/weight_free/checkpoint_key_resolver.py
+++ b/QEfficient/exporter/weight_free/checkpoint_key_resolver.py
@@ -6,8 +6,6 @@
# ----------------------------------------------------------------------------
from pathlib import Path
-from typing import Dict, List, Optional
-
import onnx_ir as ir
from torch import nn
@@ -58,7 +56,7 @@ def _collect_tied_weights(model: nn.Module) -> list[TiedWeightAlias]:
return [TiedWeightAlias(alias=alias, canonical=canonical) for alias, canonical in tied_mapping.items()]
-def _moe_weight_aliases(name: str) -> List[str]:
+def _moe_weight_aliases(name: str) -> list[str]:
"""Return equivalent checkpoint aliases for shared MoEWeights parameters."""
aliases = []
canonical = name
@@ -75,22 +73,51 @@ def _moe_weight_aliases(name: str) -> List[str]:
return aliases
-def _find_checkpoint_key(candidates: List[str], checkpoint_index: Dict[str, str], onnx_name: str) -> Optional[str]:
+def _vlm_wrapper_aliases(name: str) -> list[str]:
+ """Return aliases introduced by multimodal wrapper nesting."""
+ aliases = []
+ if name.startswith("model.model."):
+ aliases.append("model." + name[len("model.model.") :])
+ if name.startswith("model.vision_model."):
+ aliases.append("model.visual." + name[len("model.vision_model.") :])
+ if name.startswith("vision_model."):
+ aliases.append("model.visual." + name[len("vision_model.") :])
+ if name.startswith("visual."):
+ aliases.append("model.visual." + name[len("visual.") :])
+ if name.startswith("language_model."):
+ aliases.append("model.language_model." + name[len("language_model.") :])
+ if name.startswith("model.lm_head."):
+ aliases.append("lm_head." + name[len("model.lm_head.") :])
+ if name.startswith("lm_head."):
+ aliases.append("model.lm_head." + name[len("lm_head.") :])
+ if name.endswith("lm_head.weight"):
+ prefix = name[: -len("lm_head.weight")]
+ aliases.extend(
+ [
+ f"{prefix}language_model.embed_tokens.weight",
+ f"{prefix}embed_tokens.weight",
+ ]
+ )
+ return aliases
+
+
+def _find_checkpoint_key(candidates: list[str], checkpoint_index: dict[str, str], onnx_name: str) -> str | None:
"""Return the unique matching checkpoint key, or fail on ambiguous matches."""
seen = set()
matches = []
for candidate in candidates:
- if candidate in seen:
- continue
- seen.add(candidate)
- if candidate in checkpoint_index:
- matches.append(candidate)
- for alias in _moe_weight_aliases(candidate):
+ for alias in [candidate, *_vlm_wrapper_aliases(candidate)]:
if alias in seen:
continue
seen.add(alias)
if alias in checkpoint_index:
matches.append(alias)
+ for moe_alias in _moe_weight_aliases(alias):
+ if moe_alias in seen:
+ continue
+ seen.add(moe_alias)
+ if moe_alias in checkpoint_index:
+ matches.append(moe_alias)
if len(matches) > 1:
raise ValueError(
f"Ambiguous checkpoint key for ONNX initializer '{onnx_name}': matched {matches}. "
@@ -106,9 +133,9 @@ def _is_computed_initializer(name: str) -> bool:
def find_checkpoint_key(
onnx_name: str,
- checkpoint_index: Dict[str, str],
+ checkpoint_index: dict[str, str],
backbone: nn.Module,
-) -> Optional[str]:
+) -> str | None:
"""Resolve an ONNX initializer name to its safetensors checkpoint key.
Most weights match directly. The fallback rules cover wrapper prefixes,
@@ -179,15 +206,14 @@ 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
- promoted_inputs: List[WeightSpecInput] = []
+ 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)
if checkpoint_key is None:
+ if name not in model_names:
+ continue
if _is_computed_initializer(onnx_name):
continue
raise ValueError(
diff --git a/QEfficient/exporter/weight_free/export.py b/QEfficient/exporter/weight_free/export.py
index db5b8c20ae..5d4001fe79 100644
--- a/QEfficient/exporter/weight_free/export.py
+++ b/QEfficient/exporter/weight_free/export.py
@@ -35,6 +35,34 @@ def _to_meta(value: Any) -> Any:
return value
+def _iter_weight_free_configs(qeff_model):
+ model = getattr(qeff_model, "model", None)
+ nested_model = getattr(model, "model", None)
+
+ for config in (
+ getattr(model, "config", None),
+ getattr(qeff_model, "config", None),
+ getattr(getattr(model, "vision_model", None), "config", None),
+ getattr(getattr(nested_model, "vision_model", None), "config", None),
+ getattr(nested_model, "config", None),
+ ):
+ if config is not None:
+ yield config
+
+
+def _resolve_weight_free_config(qeff_model):
+ return next(_iter_weight_free_configs(qeff_model), None)
+
+
+def _resolve_weight_free_target_dtype(qeff_model) -> torch.dtype:
+ for config in _iter_weight_free_configs(qeff_model):
+ for attr in ("dtype", "torch_dtype"):
+ dtype = getattr(config, attr, None)
+ if isinstance(dtype, torch.dtype):
+ return dtype
+ return torch.float32
+
+
def _run_quantizer_for_wf(qeff_model, target_dtype: torch.dtype):
"""Finish preparing a meta-device QEfficient wrapper for weight-free tracing, in place."""
model_ref = qeff_model.hash_params.get("pretrained_model_name_or_path")
@@ -44,7 +72,8 @@ def _run_quantizer_for_wf(qeff_model, target_dtype: torch.dtype):
"Pass `pretrained_model_name_or_path=...` when constructing the QEff model manually."
)
- quant_config = getattr(qeff_model.model.config, "quantization_config", None)
+ config = _resolve_weight_free_config(qeff_model)
+ quant_config = getattr(config, "quantization_config", None)
if quant_config is not None:
# For quantized models the meta model must use the same quantized layer types as the
@@ -135,7 +164,17 @@ def _prepare_checkpoint_for_weight_free_export(
dtype_suffix = str(target_dtype).replace("torch.", "")
# TODO(wf): For different flavours of the model that expect different checkpoint weight layouts,
# we end up overriding old one. We need to add support of hashing/caching here.
- prepared_name = source_dir.name + f"-qeff-prepared-{dtype_suffix}"
+ hash_params = dict(getattr(qeff_model, "hash_params", {}) or {})
+ flavour = hash_params.get("moe_prefill_flavour")
+ if hasattr(flavour, "value"):
+ flavour = flavour.value
+ expert_parallel_suffix = ""
+ if flavour == "expert_parallel":
+ num_parallelized_experts = hash_params.get("moe_prefill_num_parallelized_experts")
+ num_pipeline_stages = hash_params.get("moe_prefill_num_pipeline_stages")
+ if num_parallelized_experts is not None and num_pipeline_stages is not None:
+ expert_parallel_suffix = f"-moe-expert-parallel-{num_parallelized_experts}x{num_pipeline_stages}"
+ prepared_name = source_dir.name + f"-qeff-prepared-{dtype_suffix}{expert_parallel_suffix}"
if QEFF_CHECKPOINT_HOME:
prepared_out = QEFF_CHECKPOINT_HOME.expanduser() / prepared_name
else:
@@ -146,6 +185,7 @@ def _prepare_checkpoint_for_weight_free_export(
src=source_dir,
out=prepared_out,
target_dtype=target_dtype,
+ hash_params=hash_params,
)
)
@@ -186,7 +226,7 @@ def export_weight_free_onnx(
tuple
Meta QEfficient model, updated ONNX transform kwargs, and cleanup callback.
"""
- target_dtype = qeff_model.model.config.dtype
+ target_dtype = _resolve_weight_free_target_dtype(qeff_model)
meta_qeff_model = _run_quantizer_for_wf(qeff_model, target_dtype)
# export_wrapper (the @export_wrapper decorator on _export) already ran
diff --git a/QEfficient/transformers/models/modeling_auto.py b/QEfficient/transformers/models/modeling_auto.py
index c9b15efe30..a26d08b765 100755
--- a/QEfficient/transformers/models/modeling_auto.py
+++ b/QEfficient/transformers/models/modeling_auto.py
@@ -10,7 +10,7 @@
import warnings
from pathlib import Path
from time import perf_counter
-from typing import List, Optional, Union
+from typing import Optional
import numpy as np
import onnx
@@ -110,18 +110,6 @@
}
-def _disable_unsupported_weight_free(kwargs: dict, qeff_auto_class_name: str) -> None:
- """Remove unsupported weight-free mode from non-CausalLM wrappers."""
-
- if not kwargs.pop("weight_free", False):
- return
-
- logger.warning(
- "weight_free=True is only supported for QEFFAutoModelForCausalLM; disabling it for %s.",
- qeff_auto_class_name,
- )
-
-
def _resolve_torch_dtype(kwargs: dict) -> None:
"""
Resolve torch_dtype in kwargs before calling from_pretrained.
@@ -286,7 +274,7 @@ def _add_retained_state_custom_io(
custom_io[compiler_output_name] = dtype
-def _filter_custom_io_for_onnx(custom_io: dict, onnx_path: Optional[Union[str, Path]]) -> dict:
+def _filter_custom_io_for_onnx(custom_io: dict, onnx_path: str | Path | None) -> dict:
"""Keep custom-IO entries that exist in the ONNX graph.
Layerwise stitched graphs may prefix I/O names (for example ``layer_0/``)
@@ -304,7 +292,7 @@ def _filter_custom_io_for_onnx(custom_io: dict, onnx_path: Optional[Union[str, P
io_names = {value.name for value in list(model.graph.input) + list(model.graph.output)}
basename_to_name = {name.rsplit("/", 1)[-1]: name for name in io_names}
- def resolve_name(name: str) -> Optional[str]:
+ def resolve_name(name: str) -> str | None:
candidates = [name]
if name.endswith("_InternalRetainedState"):
candidates.append(name[: -len("_InternalRetainedState")] + "_RetainedState")
@@ -338,7 +326,6 @@ class QEFFTransformersBase(QEFFBaseModel):
def __init__(self, model: nn.Module, **kwargs) -> None:
_configure_proxy_for_model(self, kwargs.pop("enable_proxy", False))
- _disable_unsupported_weight_free(kwargs, self.__class__.__name__)
if (
hasattr(model, "config")
@@ -377,7 +364,6 @@ def from_pretrained(cls, pretrained_model_name_or_path: str, *args, **kwargs):
QEFFTransformersBase
An instance of the specific QEFFAutoModel subclass, initialized with the pretrained weights.
"""
- _disable_unsupported_weight_free(kwargs, cls.__name__)
enable_proxy = kwargs.pop("enable_proxy", False)
if kwargs.get("attn_implementation", None) not in {None, "eager"}:
@@ -546,7 +532,6 @@ def from_pretrained(cls, pretrained_model_name_or_path, pooling=None, *args, **k
QEFFAutoModel
An instance initialized with the pretrained weights.
"""
- _disable_unsupported_weight_free(kwargs, cls.__name__)
enable_proxy = kwargs.pop("enable_proxy", False)
if kwargs.get("attn_implementation", None) not in {None, "eager"}:
@@ -584,7 +569,7 @@ def get_model_config(self) -> dict:
"""
return self.model.config.__dict__
- def export(self, export_dir: Optional[str] = None, **kwargs) -> str:
+ def export(self, export_dir: str | None = None, **kwargs) -> str:
"""
Export the model to ONNX format using ``torch.onnx.export``.
@@ -626,10 +611,10 @@ def export(self, export_dir: Optional[str] = None, **kwargs) -> str:
def compile(
self,
- onnx_path: Optional[str] = None,
- compile_dir: Optional[str] = None,
+ onnx_path: str | None = None,
+ compile_dir: str | None = None,
*,
- seq_len: Union[int, List[int]] = 32,
+ seq_len: int | list[int] = 32,
batch_size: int = 1,
num_devices: int = 1,
num_cores: int = 16, # FIXME: Make this mandatory arg
@@ -719,11 +704,11 @@ def compile(
def generate(
self,
inputs: torch.Tensor,
- device_ids: List[int] = None,
+ device_ids: list[int] = None,
runtime_ai100: bool = True,
write_io: bool = False,
- dtype: Optional[torch.dtype] = torch.float32,
- ) -> Union[torch.Tensor, np.ndarray]:
+ dtype: torch.dtype | None = torch.float32,
+ ) -> torch.Tensor | np.ndarray:
"""
Generate output by executing the compiled QPC on Cloud AI 100 hardware or using PyTorch runtime.
@@ -761,8 +746,8 @@ def generate(
def cloud_ai_100_feature_generate(
self,
inputs: torch.Tensor,
- device_ids: List[int] = None,
- dtype: Optional[torch.dtype] = torch.float32,
+ device_ids: list[int] = None,
+ dtype: torch.dtype | None = torch.float32,
) -> np.ndarray:
"""
Generate features for a batch of inputs using the Cloud AI 100 hardware runtime.
@@ -835,7 +820,7 @@ def cloud_ai_100_feature_generate(
return outputs
- def pytorch_feature_generate(self, model, inputs: Union[torch.Tensor, np.ndarray]) -> List[torch.Tensor]:
+ def pytorch_feature_generate(self, model, inputs: torch.Tensor | np.ndarray) -> list[torch.Tensor]:
"""
Generate features from a batch of inputs using the PyTorch model.
@@ -886,6 +871,13 @@ class QEFFAutoModelForSequenceClassification(QEFFTransformersBase):
_hf_auto_class = AutoModelForSequenceClassification
_pytorch_transforms = [CustomOpsTransform, TextClassificationTransform]
_onnx_transforms = []
+ _checkpoint_transforms = [
+ GptOssMxfp4ExpertDequantSplitCheckpointTransform,
+ MoEExpertStackingCheckpointTransform,
+ MoEFusedExpertSplitCheckpointTransform,
+ GraniteMoeFusedExpertSplitCheckpointTransform,
+ DtypeConversionCheckpointTransform,
+ ]
def __init__(self, model: nn.Module, **kwargs):
"""
@@ -928,7 +920,6 @@ def from_pretrained(cls, pretrained_model_name_or_path, *args, **kwargs):
QEFFAutoModelForSequenceClassification
An instance initialized with the pretrained weights.
"""
- _disable_unsupported_weight_free(kwargs, cls.__name__)
enable_proxy = kwargs.pop("enable_proxy", False)
if kwargs.get("attn_implementation", None) not in {None, "eager"}:
@@ -956,7 +947,7 @@ def get_model_config(self) -> dict:
"""
return self.model.config.__dict__
- def export(self, export_dir: Optional[str] = None, **kwargs) -> str:
+ def export(self, export_dir: str | None = None, **kwargs) -> str:
"""
Export the model to ONNX format using ``torch.onnx.export``.
@@ -998,10 +989,10 @@ def export(self, export_dir: Optional[str] = None, **kwargs) -> str:
def compile(
self,
- onnx_path: Optional[str] = None,
- compile_dir: Optional[str] = None,
+ onnx_path: str | None = None,
+ compile_dir: str | None = None,
*,
- seq_len: Union[int, List[int]] = 32,
+ seq_len: int | list[int] = 32,
batch_size: int = 1,
num_devices: int = 1,
num_cores: int = 16,
@@ -1071,7 +1062,7 @@ def compile(
def generate(
self,
inputs: torch.Tensor,
- device_ids: List[int] = None,
+ device_ids: list[int] = None,
) -> dict:
"""
Generate classification output using the Cloud AI 100 hardware runtime.
@@ -1137,6 +1128,13 @@ class QEffVisionEncoderForTextImageToTextModel(QEFFBaseModel):
KVCacheExternalModuleMapperTransform,
]
_onnx_transforms = []
+ _checkpoint_transforms = [
+ GptOssMxfp4ExpertDequantSplitCheckpointTransform,
+ MoEExpertStackingCheckpointTransform,
+ MoEFusedExpertSplitCheckpointTransform,
+ GraniteMoeFusedExpertSplitCheckpointTransform,
+ DtypeConversionCheckpointTransform,
+ ]
def __init__(self, model: nn.modules, **kwargs):
"""
@@ -1150,9 +1148,9 @@ def __init__(self, model: nn.modules, **kwargs):
Additional keyword arguments passed to the base class constructor.
"""
_configure_proxy_for_model(self, kwargs.pop("enable_proxy", False))
- _disable_unsupported_weight_free(kwargs, self.__class__.__name__)
super().__init__(model, **kwargs)
self.model = model.get_qeff_vision_encoder()
+ self.model.config = self.config
self.hash_params["qeff_auto_class"] = self.__class__.__name__
def export(self, inputs, output_names, dynamic_axes, export_dir=None, offload_pt_weights=True, **kwargs):
@@ -1186,6 +1184,7 @@ def export(self, inputs, output_names, dynamic_axes, export_dir=None, offload_pt
export_dir=export_dir,
offload_pt_weights=offload_pt_weights,
use_onnx_subfunctions=kwargs.get("use_onnx_subfunctions", False),
+ dynamo=kwargs.get("dynamo", False),
)
def compile(
@@ -1198,6 +1197,7 @@ def compile(
aic_num_cores,
custom_io,
use_onnx_subfunctions: bool = False,
+ dynamo: bool = False,
**compiler_options,
) -> str:
"""
@@ -1238,6 +1238,7 @@ def compile(
aic_num_cores=aic_num_cores,
custom_io=custom_io,
use_onnx_subfunctions=use_onnx_subfunctions,
+ dynamo=dynamo,
**compiler_options,
)
@@ -1270,14 +1271,22 @@ class QEffCausalLMForTextImageToTextModel(QEFFBaseModel):
PackQuantizedInt4ToMatMulNBitsTransform,
FP8BlockWiseDequantQwen3VLMoeTextExpertsToQwen3VLMoeTextExpertsTransform,
FP8BlockWiseDequantLinearToLinearTransform,
+ FP8DeQuantLinearToLinearTransform,
CustomOpsTransform,
KVCacheTransform,
VlmKVOffloadTransform,
SimpleDecodeMoeTransform,
]
_onnx_transforms = []
+ _checkpoint_transforms = [
+ GptOssMxfp4ExpertDequantSplitCheckpointTransform,
+ MoEExpertStackingCheckpointTransform,
+ MoEFusedExpertSplitCheckpointTransform,
+ GraniteMoeFusedExpertSplitCheckpointTransform,
+ DtypeConversionCheckpointTransform,
+ ]
- def __init__(self, model, qaic_config: Optional[dict] = None, **kwargs):
+ def __init__(self, model, qaic_config: dict | None = None, **kwargs):
"""
Initializes the language decoder component for multimodal models.
@@ -1292,22 +1301,33 @@ def __init__(self, model, qaic_config: Optional[dict] = None, **kwargs):
Additional keyword arguments passed to the base class constructor.
"""
_configure_proxy_for_model(self, kwargs.pop("enable_proxy", False))
- _disable_unsupported_weight_free(kwargs, self.__class__.__name__)
+ if kwargs.pop("fp8_retain_weights", False):
+ self._pytorch_transforms = [
+ t
+ for t in self._pytorch_transforms
+ if t
+ not in (
+ FP8DeQuantLinearToLinearTransform,
+ FP8BlockWiseDequantLinearToLinearTransform,
+ FP8BlockWiseDequantQwen3VLMoeTextExpertsToQwen3VLMoeTextExpertsTransform,
+ )
+ ]
super().__init__(model, **kwargs)
self.model = model.get_qeff_language_decoder()
+ self.model.config = self.config
self.model.qaic_config = qaic_config
self.hash_params["qeff_auto_class"] = self.__class__.__name__
self.continuous_batching = False
if qaic_config:
if mla_absorption := qaic_config.get("mla_absorption", None):
self.hash_params["mla_absorption"] = mla_absorption
- setattr(self.model.language_model, "mla_absorption", mla_absorption)
+ self.model.language_model.mla_absorption = mla_absorption
def __update_prefill_transform(
self,
- enable: Optional[bool] = True,
- enable_chunking: Optional[bool] = False,
- retain_full_kv: Optional[bool] = False,
+ enable: bool | None = True,
+ enable_chunking: bool | None = False,
+ retain_full_kv: bool | None = False,
):
if enable:
if enable_chunking:
@@ -1328,10 +1348,10 @@ def export(
dynamic_axes,
export_dir=None,
offload_pt_weights=True,
- prefill_seq_len: Optional[int] = None,
+ prefill_seq_len: int | None = None,
prefill_only: bool = False,
enable_chunking: bool = False,
- kv_cache_prefix: Optional[str] = None,
+ kv_cache_prefix: str | None = None,
**kwargs,
):
"""
@@ -1396,6 +1416,7 @@ def export(
export_dir=export_dir,
offload_pt_weights=offload_pt_weights,
use_onnx_subfunctions=kwargs.get("use_onnx_subfunctions", False),
+ dynamo=kwargs.get("dynamo", False),
)
def compile(
@@ -1408,6 +1429,7 @@ def compile(
aic_num_cores,
custom_io,
use_onnx_subfunctions: bool = False,
+ dynamo: bool = False,
**compiler_options,
) -> str:
"""
@@ -1448,6 +1470,7 @@ def compile(
aic_num_cores=aic_num_cores,
custom_io=custom_io,
use_onnx_subfunctions=use_onnx_subfunctions,
+ dynamo=dynamo,
**compiler_options,
)
@@ -1481,7 +1504,7 @@ def __init__(
self,
model: nn.Module,
continuous_batching: bool = False,
- qaic_config: Optional[dict] = None,
+ qaic_config: dict | None = None,
**kwargs,
):
"""
@@ -1496,7 +1519,6 @@ def __init__(
**kwargs :
Additional keyword arguments.
"""
- _disable_unsupported_weight_free(kwargs, self.__class__.__name__)
if kwargs.pop("full_batch_size", None):
continuous_batching = True
warnings.warn(
@@ -1505,6 +1527,7 @@ def __init__(
self.model = model
self.config = model.config
self._pretrained_model_name_or_path = kwargs.get("pretrained_model_name_or_path", None)
+ self._weight_free = kwargs.get("weight_free", False)
self.vision_model = QEffVisionEncoderForTextImageToTextModel(model, **kwargs)
self.lang_model = QEffCausalLMForTextImageToTextModel(model, qaic_config=qaic_config, **kwargs)
@@ -1522,7 +1545,13 @@ def __init__(
self.lang_model.model, _ = SamplerTransform.apply(self.lang_model.model, qaic_config, **kwargs)
@classmethod
- def from_pretrained(cls, pretrained_model_name_or_path: str, qaic_config: Optional[dict] = None, **kwargs):
+ def from_pretrained(
+ cls,
+ pretrained_model_name_or_path: str,
+ qaic_config: dict | None = None,
+ weight_free: bool = False,
+ **kwargs,
+ ):
"""
Load a QEfficient multimodal model for dual QPC from a pretrained HuggingFace model or local path.
@@ -1540,7 +1569,6 @@ def from_pretrained(cls, pretrained_model_name_or_path: str, qaic_config: Option
_QEffAutoModelForImageTextToTextDualQPC
An instance initialized with the pretrained weights.
"""
- _disable_unsupported_weight_free(kwargs, cls.__name__)
enable_proxy = kwargs.pop("enable_proxy", False)
if kwargs.get("attn_implementation", None) not in {None, "eager"}:
@@ -1559,7 +1587,10 @@ def from_pretrained(cls, pretrained_model_name_or_path: str, qaic_config: Option
)
_resolve_torch_dtype(kwargs)
- model = cls._hf_auto_class.from_pretrained(pretrained_model_name_or_path, **kwargs)
+ if weight_free:
+ model = _build_meta_model(cls._hf_auto_class, pretrained_model_name_or_path, kwargs)
+ else:
+ model = cls._hf_auto_class.from_pretrained(pretrained_model_name_or_path, **kwargs)
kwargs.update({"enable_proxy": enable_proxy} if enable_proxy else {})
@@ -1567,6 +1598,7 @@ def from_pretrained(cls, pretrained_model_name_or_path: str, qaic_config: Option
model,
pretrained_model_name_or_path=pretrained_model_name_or_path,
qaic_config=qaic_config,
+ weight_free=weight_free,
**kwargs,
)
@@ -1582,11 +1614,16 @@ def onnx_path(self):
"""
return [self.vision_model.onnx_path, self.lang_model.onnx_path]
+ @property
+ def weight_spec_path(self):
+ """Get the weight-spec paths for the vision and language components."""
+ return [self.vision_model.weight_spec_path, self.lang_model.weight_spec_path]
+
def __update_prefill_transform(
self,
- enable: Optional[bool] = True,
- enable_chunking: Optional[bool] = False,
- retain_full_kv: Optional[bool] = False,
+ enable: bool | None = True,
+ enable_chunking: bool | None = False,
+ retain_full_kv: bool | None = False,
):
if enable:
self.model, tf = PrefillOnlyExternalModuleMapperTransform.apply(self.model)
@@ -1604,18 +1641,19 @@ def __update_prefill_transform(
def export(
self,
- export_dir: Optional[str] = None,
+ export_dir: str | None = None,
use_onnx_subfunctions: bool = False,
- skip_vision: Optional[bool] = False,
- skip_lang: Optional[bool] = False,
- prefill_seq_len: Optional[int] = None,
+ skip_vision: bool | None = False,
+ skip_lang: bool | None = False,
+ prefill_seq_len: int | None = None,
prefill_only: bool = False,
enable_chunking: bool = False,
num_cores: int = constants.DEFAULT_AIC_NUM_CORES,
layerwise: bool = False,
layerwise_window_size: int = 1,
- kv_cache_prefix: Optional[str] = None,
- offload_pt_weights: Optional[bool] = None,
+ kv_cache_prefix: str | None = None,
+ offload_pt_weights: bool | None = None,
+ dynamo: bool = False,
**kwargs,
) -> str:
"""
@@ -1640,6 +1678,10 @@ def export(
"""
layerwise_cache_probe = kwargs.pop("_layerwise_cache_probe", False)
reject_legacy_moe_prefill_packed_chunk_size(kwargs)
+ dynamo = dynamo or self._weight_free
+ use_onnx_subfunctions = use_onnx_subfunctions or self._weight_free
+ if layerwise and dynamo:
+ raise NotImplementedError("Dynamo export is not supported for layerwise VLM export.")
if layerwise:
return self._run_layerwise_export(
export_dir=export_dir,
@@ -1654,7 +1696,7 @@ def export(
**kwargs,
)
bs: int = constants.ONNX_EXPORT_EXAMPLE_BATCH_SIZE
- seq_len: int = constants.ONNX_EXPORT_EXAMPLE_SEQ_LEN
+ seq_len: int = prefill_seq_len if prefill_seq_len else constants.ONNX_EXPORT_EXAMPLE_SEQ_LEN
qaic_config = kwargs.get("qaic_config", getattr(self.lang_model.model, "qaic_config", None))
# TODO: move this to a DA Serving utility class
if self.model.config.model_type in SPECIALIZED_DISAGG_SERVING_MODEL_ARCH:
@@ -1662,7 +1704,7 @@ def export(
self.__update_prefill_transform(enable=True, enable_chunking=enable_chunking)
else:
self.__update_prefill_transform(False, retain_full_kv=kwargs.get("retain_full_kv", False))
- onnx_kwargs = {"prefill_seq_len": seq_len, "batch_size": bs}
+ onnx_kwargs = {"prefill_seq_len": seq_len, "batch_size": bs, "ctx_len": kwargs.get("ctx_len", seq_len)}
dynamic_axes_kwargs = {
"kv_offload": True,
"continuous_batching": self.continuous_batching,
@@ -1721,6 +1763,7 @@ def export(
export_dir=export_dir,
offload_pt_weights=False,
use_onnx_subfunctions=use_onnx_subfunctions,
+ dynamo=dynamo,
)
# TODO: remove the current pt weight offload capability once CustomLoader is in place
@@ -1745,16 +1788,17 @@ def export(
qaic_config=qaic_config,
_layerwise_cache_probe=layerwise_cache_probe,
kv_cache_prefix=kv_cache_prefix,
+ dynamo=dynamo,
)
return self.onnx_path
def transform(
self,
- ctx_len: Optional[int] = None,
- seq_len: Optional[int] = None,
- bs: Optional[int] = 1,
+ ctx_len: int | None = None,
+ seq_len: int | None = None,
+ bs: int | None = 1,
num_devices: int = 1,
- qaic_config: Optional[dict] = None,
+ qaic_config: dict | None = None,
**compiler_options,
):
self.vision_model.transform(
@@ -1903,32 +1947,34 @@ def _run_layerwise_compile(
def compile(
self,
- img_size: Optional[int] = None,
- vision_onnx_path: Optional[str] = None,
- lang_onnx_path: Optional[str] = None,
- compile_dir: Optional[str] = None,
+ img_size: int | None = None,
+ vision_onnx_path: str | None = None,
+ lang_onnx_path: str | None = None,
+ compile_dir: str | None = None,
*,
- prefill_seq_len: Optional[int] = None,
- comp_ctx_lengths_prefill: Optional[List[int]] = None,
- comp_ctx_lengths_decode: Optional[List[int]] = None,
- ctx_len: Optional[int] = None,
+ prefill_seq_len: int | None = None,
+ comp_ctx_lengths_prefill: list[int] | None = None,
+ comp_ctx_lengths_decode: list[int] | None = None,
+ ctx_len: int | None = None,
batch_size: int = 1,
- full_batch_size: Optional[int] = None,
- kv_cache_batch_size: Optional[int] = None,
+ full_batch_size: int | None = None,
+ kv_cache_batch_size: int | None = None,
num_devices: int = 1,
num_cores: int = 16, # FIXME: Make this mandatory arg
mxfp6_matmul: bool = False,
mxint8_kv_cache: bool = False,
- skip_vision: Optional[bool] = False,
- skip_lang: Optional[bool] = False,
+ skip_vision: bool | None = False,
+ skip_lang: bool | None = False,
use_onnx_subfunctions: bool = False,
prefill_only=None,
- offload_pt_weights: Optional[bool] = None,
+ offload_pt_weights: bool | None = None,
enable_chunking=False,
- qaic_config: Optional[dict] = None,
+ qaic_config: dict | None = None,
layerwise: bool = False,
layerwise_window_size: int = 1,
kv_cache_prefix: Optional[str] = None,
+ moe_prefill_packed_chunk_size: int = constants.MOE_PREFILL_PACKED_CHUNK_SIZE,
+ dynamo: bool = False,
**compiler_options,
) -> str:
"""
@@ -1991,6 +2037,10 @@ def compile(
raise ValueError("Expected at least one of 'skip_lang' or 'skip_vision' to be False")
reject_legacy_moe_prefill_packed_chunk_size(compiler_options)
_ignore_public_mdp_ts_num_devices(compiler_options)
+ dynamo = dynamo or self._weight_free
+ use_onnx_subfunctions = use_onnx_subfunctions or self._weight_free
+ if layerwise and dynamo:
+ raise NotImplementedError("Dynamo export is not supported for layerwise VLM compilation.")
if layerwise:
if skip_lang and not skip_vision:
@@ -2143,11 +2193,13 @@ def compile(
prefill_only=prefill_only,
enable_chunking=enable_chunking,
prefill_seq_len=prefill_seq_len,
+ ctx_len=ctx_len,
num_cores=num_cores,
qaic_config=qaic_config,
_layerwise_cache_probe=layerwise_cache_probe,
kv_cache_prefix=kv_cache_prefix,
offload_pt_weights=offload_pt_weights,
+ dynamo=dynamo,
)
if layerwise_cache_probe:
return self.lang_model.onnx_path
@@ -2254,21 +2306,21 @@ def compile(
def generate(
self,
- inputs: Optional[torch.Tensor] = None,
- tokenizer: Union[PreTrainedTokenizerFast, PreTrainedTokenizer] = None,
- processor: Optional[AutoImageProcessor] = None,
- images: List[str] = None,
- prompts: List[str] = None,
- streamer: Optional[TextStreamer] = None,
- device_ids: List[int] = None,
+ inputs: torch.Tensor | None = None,
+ tokenizer: PreTrainedTokenizerFast | PreTrainedTokenizer = None,
+ processor: AutoImageProcessor | None = None,
+ images: list[str] = None,
+ prompts: list[str] = None,
+ streamer: TextStreamer | None = None,
+ device_ids: list[int] = None,
runtime_ai100: bool = True,
- generation_len: Optional[int] = None,
- image_height: Optional[int] = None,
- image_width: Optional[int] = None,
- multi_specs: Optional[bool] = None,
- num_frames: Optional[int] = None,
+ generation_len: int | None = None,
+ image_height: int | None = None,
+ image_width: int | None = None,
+ multi_specs: bool | None = None,
+ num_frames: int | None = None,
**kwargs,
- ) -> Union[torch.Tensor, np.ndarray]:
+ ) -> torch.Tensor | np.ndarray:
"""
Generates output by executing the compiled QPC(s) on Cloud AI 100 Hardware cards.
@@ -2353,9 +2405,9 @@ def generate(
def kv_offload_generate(
self,
- inputs: List[str] = None,
- streamer: Optional[TextStreamer] = None,
- device_ids: List[int] = None,
+ inputs: list[str] = None,
+ streamer: TextStreamer | None = None,
+ device_ids: list[int] = None,
generation_len: int = None,
):
"""
@@ -2698,7 +2750,7 @@ class _QEFFAutoModelForImageTextToTextSingleQPC(QEFFTransformersBase, Multimodal
def __init__(
self,
model: nn.Module,
- qaic_config: Optional[dict] = None,
+ qaic_config: dict | None = None,
**kwargs,
):
"""
@@ -2751,7 +2803,7 @@ def __init__(
def from_pretrained(
cls,
pretrained_model_name_or_path,
- qaic_config: Optional[dict] = None,
+ qaic_config: dict | None = None,
*args,
**kwargs,
):
@@ -2775,7 +2827,6 @@ def from_pretrained(
_QEFFAutoModelForImageTextToTextSingleQPC
An instance initialized with the pretrained weights.
"""
- _disable_unsupported_weight_free(kwargs, cls.__name__)
enable_proxy = kwargs.pop("enable_proxy", False)
if kwargs.get("attn_implementation", None) not in {None, "eager"}:
@@ -2805,9 +2856,9 @@ def from_pretrained(
def __update_prefill_transform(
self,
- enable: Optional[bool] = True,
- enable_chunking: Optional[bool] = False,
- retain_full_kv: Optional[bool] = False,
+ enable: bool | None = True,
+ enable_chunking: bool | None = False,
+ retain_full_kv: bool | None = False,
):
if enable:
if enable_chunking:
@@ -2823,12 +2874,12 @@ def __update_prefill_transform(
def export(
self,
- export_dir: Optional[str] = None,
+ export_dir: str | None = None,
use_onnx_subfunctions: bool = False,
- prefill_seq_len: Optional[int] = None,
+ prefill_seq_len: int | None = None,
prefill_only: bool = False,
enable_chunking: bool = False,
- kv_cache_prefix: Optional[str] = None,
+ kv_cache_prefix: str | None = None,
**kwargs,
) -> str:
"""
@@ -2870,29 +2921,30 @@ def export(
dynamic_axes=dynamic_axes,
export_dir=export_dir,
use_onnx_subfunctions=use_onnx_subfunctions,
+ dynamo=kwargs.get("dynamo", False),
)
def compile(
self,
- onnx_path: Optional[str] = None,
- img_size: Optional[int] = None,
- compile_dir: Optional[str] = None,
+ onnx_path: str | None = None,
+ img_size: int | None = None,
+ compile_dir: str | None = None,
*,
- prefill_seq_len: Optional[int] = None,
- ctx_len: Optional[int] = None,
- comp_ctx_lengths_prefill: Optional[List[int]] = None,
- comp_ctx_lengths_decode: Optional[List[int]] = None,
+ prefill_seq_len: int | None = None,
+ ctx_len: int | None = None,
+ comp_ctx_lengths_prefill: list[int] | None = None,
+ comp_ctx_lengths_decode: list[int] | None = None,
batch_size: int = 1,
- full_batch_size: Optional[int] = None,
- kv_cache_batch_size: Optional[int] = None,
+ full_batch_size: int | None = None,
+ kv_cache_batch_size: int | None = None,
num_devices: int = 1,
num_cores: int = 16, # FIXME: Make this mandatory arg
mxfp6_matmul: bool = False,
mxint8_kv_cache: bool = False,
- num_speculative_tokens: Optional[int] = None,
+ num_speculative_tokens: int | None = None,
use_onnx_subfunctions: bool = False,
- qaic_config: Optional[dict] = None,
- kv_cache_prefix: Optional[str] = None,
+ qaic_config: dict | None = None,
+ kv_cache_prefix: str | None = None,
**compiler_options,
) -> str:
"""
@@ -3038,12 +3090,12 @@ def get_onnx_dynamic_axes(self):
def generate(
self,
inputs: torch.Tensor,
- streamer: Optional[TextStreamer] = None,
- device_ids: List[int] = None,
+ streamer: TextStreamer | None = None,
+ device_ids: list[int] = None,
runtime_ai100: bool = True,
- generation_len: Optional[int] = None,
+ generation_len: int | None = None,
write_io: bool = False,
- ) -> Union[torch.Tensor, np.ndarray]:
+ ) -> torch.Tensor | np.ndarray:
"""
Generates output by executing the compiled single QPC on Cloud AI 100 Hardware cards.
@@ -3085,10 +3137,10 @@ def generate(
def cloud_ai_100_generate(
self,
inputs: torch.Tensor,
- device_ids: List[int],
+ device_ids: list[int],
enable_debug_logs: bool = False,
generation_len: int = None,
- streamer: Optional[TextStreamer] = None,
+ streamer: TextStreamer | None = None,
) -> np.ndarray:
"""
Performs generation for multimodal models using a single QPC on Cloud AI 100 hardware.
@@ -3354,9 +3406,9 @@ class QEFFAutoModelForImageTextToText:
def __new__(
self,
model: nn.Module,
- kv_offload: Optional[bool] = True,
+ kv_offload: bool | None = True,
continuous_batching: bool = False,
- qaic_config: Optional[dict] = None,
+ qaic_config: dict | None = None,
**kwargs,
):
"""
@@ -3378,7 +3430,8 @@ def __new__(
Union[_QEffAutoModelForImageTextToTextDualQPC, _QEFFAutoModelForImageTextToTextSingleQPC]
The wrapped model instance, configured for either dual or single QPC.
"""
- _disable_unsupported_weight_free(kwargs, self.__name__)
+ if kwargs.get("weight_free", False) and kv_offload is not True:
+ raise NotImplementedError("weight_free=True for VLM is supported only with kv_offload=True.")
if kv_offload:
return _QEffAutoModelForImageTextToTextDualQPC(
model, continuous_batching, qaic_config=qaic_config, **kwargs
@@ -3391,10 +3444,11 @@ def __new__(
def from_pretrained(
cls,
pretrained_model_name_or_path: str,
- kv_offload: Optional[bool] = None,
+ kv_offload: bool | None = None,
continuous_batching: bool = False,
- qaic_config: Optional[dict] = None,
+ qaic_config: dict | None = None,
layerwise: bool = False,
+ weight_free: bool = False,
**kwargs,
):
"""
@@ -3426,7 +3480,14 @@ def from_pretrained(
NotImplementedError
If `continuous_batching` is provided as True.
"""
- _disable_unsupported_weight_free(kwargs, cls.__name__)
+ if layerwise and weight_free:
+ raise ValueError("`layerwise=True` and `weight_free=True` are mutually exclusive for VLM.")
+ if weight_free:
+ validate_dynamo_export_requirements("weight_free=True")
+ if kv_offload is False:
+ raise NotImplementedError("weight_free=True for VLM is supported only with kv_offload=True.")
+ kv_offload = True
+
enable_proxy = kwargs.pop("enable_proxy", False)
# TODO: add a check to see if kv_offload is allowed for given model by loading the config and checking architecture or type of config here.
@@ -3448,24 +3509,27 @@ def from_pretrained(
)
_resolve_torch_dtype(kwargs)
- if layerwise:
+ fp8_retain_weights = kwargs.pop("fp8_retain_weights", False)
+ if layerwise or weight_free:
# Layer-wise mode: build the outer model on the meta device so the
# caller's ``from_pretrained`` does not pull the full checkpoint
- # into RAM. compile()/export() rebuilds a real per-window model
- # internally via the layer-wise driver, so the outer instance is
- # only used as a config holder.
+ # into RAM. Weight-free export later supplies checkpoint tensors
+ # through the generated weight specification.
model = _build_meta_model(cls._hf_auto_class, pretrained_model_name_or_path, kwargs)
else:
model = cls._hf_auto_class.from_pretrained(pretrained_model_name_or_path, **kwargs)
kwargs.update({"enable_proxy": enable_proxy} if enable_proxy else {})
+ if fp8_retain_weights:
+ kwargs["fp8_retain_weights"] = fp8_retain_weights
instance = cls(
model,
kv_offload=kv_offload,
continuous_batching=continuous_batching,
pretrained_model_name_or_path=pretrained_model_name_or_path,
qaic_config=qaic_config,
+ weight_free=weight_free,
**kwargs,
)
# Mark the wrapper so its compile() can default ``layerwise=True`` if
@@ -3508,6 +3572,7 @@ class QEFFAutoModelForCausalLM(QEFFBaseModel):
AwqToMatmulNbitsTransform,
GPTQToMatmulNbitsTransform,
FP8DeQuantLinearToLinearTransform,
+ FP8BlockWiseDequantLinearToLinearTransform,
PackQuantizedInt4ToMatMulNBitsTransform,
Mxfp4GptOssExpertDequantizeTransform,
CustomOpsTransform,
@@ -3528,9 +3593,9 @@ class QEFFAutoModelForCausalLM(QEFFBaseModel):
def prefill(
self,
- enable: Optional[bool] = True,
- enable_chunking: Optional[bool] = False,
- retain_full_kv: Optional[bool] = False,
+ enable: bool | None = True,
+ enable_chunking: bool | None = False,
+ retain_full_kv: bool | None = False,
):
if enable:
self.model, tf = PrefillOnlyExternalModuleMapperTransform.apply(self.model)
@@ -3548,9 +3613,9 @@ def prefill(
def __update_prefill_transform(
self,
- enable: Optional[bool] = True,
- enable_chunking: Optional[bool] = False,
- retain_full_kv: Optional[bool] = False,
+ enable: bool | None = True,
+ enable_chunking: bool | None = False,
+ retain_full_kv: bool | None = False,
):
if enable:
self.model, tf = PrefillOnlyExternalModuleMapperTransform.apply(self.model)
@@ -3570,8 +3635,8 @@ def __init__(
self,
model: nn.Module,
continuous_batching: bool = False,
- qaic_config: Optional[dict] = None,
- max_seq_len_cached: Optional[int] = None,
+ qaic_config: dict | None = None,
+ max_seq_len_cached: int | None = None,
**kwargs,
):
"""
@@ -3619,10 +3684,16 @@ def __init__(
logger.warning(
"Please use `from_pretrained` method to load quantized models, might give unexpected results"
)
+ if kwargs.pop("fp8_retain_weights", False):
+ self._pytorch_transforms = [
+ t
+ for t in self._pytorch_transforms
+ if t not in (FP8DeQuantLinearToLinearTransform, FP8BlockWiseDequantLinearToLinearTransform)
+ ]
# Set use_cache=True to get KV values as output during ONNX export
model.config.use_cache = True
- setattr(model.config, "max_seq_len_cached", max_seq_len_cached)
+ model.config.max_seq_len_cached = max_seq_len_cached
super().__init__(model, qaic_config=qaic_config, **kwargs)
self.num_layers = model.config.num_hidden_layers
self.continuous_batching = continuous_batching
@@ -3637,7 +3708,7 @@ def __init__(
self.ccl_enabled = qaic_config.get("ccl_enabled", False)
if mla_absorption := qaic_config.get("mla_absorption", None):
self.hash_params["mla_absorption"] = mla_absorption
- setattr(self.model, "mla_absorption", mla_absorption)
+ self.model.mla_absorption = mla_absorption
self.comp_ctx_lengths_prefill, self.comp_ctx_lengths_decode = None, None
self.hash_params["max_seq_len_cached"] = max_seq_len_cached
@@ -3660,8 +3731,8 @@ def from_pretrained(
cls,
pretrained_model_name_or_path,
continuous_batching: bool = False,
- qaic_config: Optional[dict] = None,
- max_seq_len_cached: Optional[int] = None,
+ qaic_config: dict | None = None,
+ max_seq_len_cached: int | None = None,
layerwise: bool = False,
weight_free: bool = False,
*args,
@@ -3752,6 +3823,7 @@ def from_pretrained(
}
)
+ fp8_retain_weights = kwargs.pop("fp8_retain_weights", False)
_resolve_torch_dtype(kwargs)
if layerwise:
warnings.warn(
@@ -3778,6 +3850,8 @@ def from_pretrained(
# This is support models that should be classified to in a different auto class but transformers load them via this class
kwargs.update({"enable_proxy": enable_proxy} if enable_proxy else {})
+ if fp8_retain_weights:
+ kwargs["fp8_retain_weights"] = fp8_retain_weights
if model.__class__.__name__ in MISCLASSIFIED_CAUSAL_LM_TO_QEFF_AUTO_CLASS_MAP:
return MISCLASSIFIED_CAUSAL_LM_TO_QEFF_AUTO_CLASS_MAP[model.__class__.__name__](
model,
@@ -3846,11 +3920,7 @@ def handle_gpt_oss_env_variable_legacy_burden(self, prefill_seq_len: Optional[in
self.hash_params["NUM_Q_BLOCKS"] = num_q_blocks
self.hash_params["NUM_FFN_BLOCKS"] = num_ffn_blocks
self.hash_params["ENABLE_OPT_SWA"] = os.environ.get("ENABLE_OPT_SWA", "0")
- return (
- min_seq_len
- if min_seq_len > constants.ONNX_EXPORT_EXAMPLE_SEQ_LEN
- else constants.ONNX_EXPORT_EXAMPLE_SEQ_LEN
- )
+ return max(constants.ONNX_EXPORT_EXAMPLE_SEQ_LEN, min_seq_len)
def _run_layerwise(self, *, final_compile: bool, layerwise_window_size: int, **forward_kwargs):
"""Drive the layer-wise export/compile loop for CausalLM models."""
@@ -3887,13 +3957,13 @@ def _factory(model_id, config):
def export(
self,
- export_dir: Optional[str] = None,
- prefill_only: Optional[bool] = False,
- prefill_seq_len: Optional[int] = None,
+ export_dir: str | None = None,
+ prefill_only: bool | None = False,
+ prefill_seq_len: int | None = None,
num_cores: int = constants.DEFAULT_AIC_NUM_CORES,
layerwise: bool = False,
layerwise_window_size: int = 1,
- kv_cache_prefix: Optional[str] = None,
+ kv_cache_prefix: str | None = None,
dynamo: bool = False,
**kwargs,
) -> str:
@@ -4268,10 +4338,10 @@ def build_prefill_specialization(
self,
prefill_seq_len: int = 32,
ctx_len: int = 128,
- comp_ctx_lengths: Optional[int] = None,
+ comp_ctx_lengths: int | None = None,
batch_size: int = 1,
- kv_cache_batch_size: Optional[int] = None,
- full_batch_size: Optional[int] = None,
+ kv_cache_batch_size: int | None = None,
+ full_batch_size: int | None = None,
**kwargs,
):
"""
@@ -4333,11 +4403,11 @@ def build_decode_specialization(
self,
prefill_seq_len: int = 32,
ctx_len: int = 128,
- comp_ctx_lengths: Optional[int] = None,
+ comp_ctx_lengths: int | None = None,
batch_size: int = 1,
- kv_cache_batch_size: Optional[int] = None,
- full_batch_size: Optional[int] = None,
- num_speculative_tokens: Optional[int] = None,
+ kv_cache_batch_size: int | None = None,
+ full_batch_size: int | None = None,
+ num_speculative_tokens: int | None = None,
**kwargs,
):
"""
@@ -4394,29 +4464,29 @@ def build_decode_specialization(
def compile(
self,
- onnx_path: Optional[str] = None,
- compile_dir: Optional[str] = None,
+ onnx_path: str | None = None,
+ compile_dir: str | None = None,
*,
prefill_seq_len: int = 32,
ctx_len: int = 128,
- comp_ctx_lengths_prefill: Optional[List[int]] = None,
- comp_ctx_lengths_decode: Optional[List[int]] = None,
+ comp_ctx_lengths_prefill: list[int] | None = None,
+ comp_ctx_lengths_decode: list[int] | None = None,
batch_size: int = 1,
- full_batch_size: Optional[int] = None,
- kv_cache_batch_size: Optional[int] = None,
+ full_batch_size: int | None = None,
+ kv_cache_batch_size: int | None = None,
num_devices: int = 1,
num_cores: int = 16, # FIXME: Make this mandatory arg
mxfp6_matmul: bool = False,
mxint8_kv_cache: bool = False,
- num_speculative_tokens: Optional[Union[int, List[int]]] = None,
- prefill_only: Optional[bool] = None,
+ num_speculative_tokens: int | list[int] | None = None,
+ prefill_only: bool | None = None,
use_onnx_subfunctions: bool = False,
offload_pt_weights: Optional[bool] = True,
enable_chunking: Optional[bool] = False,
retain_full_kv: Optional[bool] = None,
layerwise: bool = False,
layerwise_window_size: int = 1,
- kv_cache_prefix: Optional[str] = None,
+ kv_cache_prefix: str | None = None,
**compiler_options,
) -> str:
"""
@@ -4597,11 +4667,6 @@ def compile(
if prefill_only is not None and not isinstance(prefill_only, bool):
raise TypeError("`prefill_only` must be a boolean.")
- if self._weight_free and (prefill_only is True or prefill_seq_len == 1):
- raise NotImplementedError(
- "weight_free=True is not supported with disaggregated compile (prefill_only=True or prefill_seq_len=1)."
- )
-
_decode_ks = (
sorted(set(num_speculative_tokens))
if isinstance(num_speculative_tokens, (list, tuple))
@@ -4646,7 +4711,7 @@ def compile(
if self.comp_ctx_lengths_prefill is not None or self.comp_ctx_lengths_decode is not None:
ccl_lengths = self.comp_ctx_lengths_decode if prefill_seq_len == 1 else self.comp_ctx_lengths_prefill
# Adding elements from self.comp_ctx_lengths_prefill to prefill_specialization
- for i in range(0, len(ccl_lengths)):
+ for i in range(len(ccl_lengths)):
specializations.append(
self.build_prefill_specialization(
prefill_seq_len=prefill_seq_len,
@@ -4700,7 +4765,7 @@ def compile(
elif self.comp_ctx_lengths_decode is not None:
# CCL loop (non-TLM)
- for i in range(0, len(self.comp_ctx_lengths_decode)):
+ for i in range(len(self.comp_ctx_lengths_decode)):
decode_spec = self.build_decode_specialization(
prefill_seq_len=prefill_seq_len,
ctx_len=ctx_len,
@@ -4794,9 +4859,9 @@ def compile(
# FIXME: Update this method to match with transformers AutoModelForCausalLM.generate
def generate(
self,
- tokenizer: Union[PreTrainedTokenizerFast, PreTrainedTokenizer],
- prompts: List[str],
- device_id: List[int] = None,
+ tokenizer: PreTrainedTokenizerFast | PreTrainedTokenizer,
+ prompts: list[str],
+ device_id: list[int] = None,
runtime_ai100: bool = True,
**kwargs,
):
@@ -4857,7 +4922,7 @@ def generate(
else:
raise NotImplementedError("Only AI_100 runtime is supported right now via generate API")
- def check_and_get_num_speculative_tokens(self, num_speculative_tokens: Optional[int], prefill_seq_len: int):
+ def check_and_get_num_speculative_tokens(self, num_speculative_tokens: int | None, prefill_seq_len: int):
"""
Validates and retrieves the number of speculative tokens for TLM models.
@@ -4985,7 +5050,7 @@ def get_model_config(self) -> dict:
"""
return self.model.config.__dict__
- def export(self, export_dir: Optional[str] = None, **kwargs) -> str:
+ def export(self, export_dir: str | None = None, **kwargs) -> str:
"""
Export the model to ONNX format using ``torch.onnx.export``.
@@ -5018,20 +5083,20 @@ def export(self, export_dir: Optional[str] = None, **kwargs) -> str:
def compile(
self,
- onnx_path: Optional[str] = None,
- compile_dir: Optional[str] = None,
+ onnx_path: str | None = None,
+ compile_dir: str | None = None,
*,
- prefill_seq_len: Optional[int] = 1,
- encoder_ctx_len: Optional[int] = None,
+ prefill_seq_len: int | None = 1,
+ encoder_ctx_len: int | None = None,
ctx_len: int = 150,
- full_batch_size: Optional[int] = None,
- kv_cache_batch_size: Optional[int] = None,
+ full_batch_size: int | None = None,
+ kv_cache_batch_size: int | None = None,
batch_size: int = 1,
num_devices: int = 1,
num_cores: int = 16, # FIXME: Make this mandatory arg
mxfp6_matmul: bool = False,
mxint8_kv_cache: bool = False,
- num_speculative_tokens: Optional[int] = None,
+ num_speculative_tokens: int | None = None,
use_onnx_subfunctions: bool = False,
**compiler_options,
) -> str:
@@ -5151,10 +5216,10 @@ def generate(
self,
inputs: torch.Tensor,
generation_len: int,
- streamer: Optional[TextStreamer] = None,
- device_ids: List[int] = None,
+ streamer: TextStreamer | None = None,
+ device_ids: list[int] = None,
write_io: bool = False,
- ) -> Union[torch.Tensor, np.ndarray]:
+ ) -> torch.Tensor | np.ndarray:
"""
Generate output until ``<|endoftext|>`` token or `generation_len` is reached,
by executing the compiled QPC on Cloud AI 100 hardware.
@@ -5348,7 +5413,6 @@ def from_pretrained(cls, pretrained_model_name_or_path, pooling=None, *args, **k
# You can now execute the model
out = model.generate(processor,inputs=input_audio)
"""
- _disable_unsupported_weight_free(kwargs, cls.__name__)
enable_proxy = kwargs.pop("enable_proxy", False)
if kwargs.get("attn_implementation", None) not in {None, "eager"}:
logger.warning('Updating attn_implementation="eager"')
@@ -5377,7 +5441,7 @@ def from_pretrained(cls, pretrained_model_name_or_path, pooling=None, *args, **k
def get_model_config(self) -> dict:
return self.model.config.__dict__
- def export(self, export_dir: Optional[str] = None, **kwargs) -> str:
+ def export(self, export_dir: str | None = None, **kwargs) -> str:
"""
Exports the model to ``ONNX`` format using ``torch.onnx.export``.
@@ -5410,10 +5474,10 @@ def export(self, export_dir: Optional[str] = None, **kwargs) -> str:
def compile(
self,
- onnx_path: Optional[str] = None,
- compile_dir: Optional[str] = None,
+ onnx_path: str | None = None,
+ compile_dir: str | None = None,
*,
- seq_len: Union[int, List[int]] = 480000,
+ seq_len: int | list[int] = 480000,
batch_size: int = 1,
num_devices: int = 1,
num_cores: int = 16, # FIXME: Make this mandatory arg
@@ -5478,10 +5542,10 @@ def generate(
self,
processor,
inputs: torch.Tensor,
- device_ids: List[int] = None,
+ device_ids: list[int] = None,
runtime_ai100: bool = True,
write_io: bool = False,
- ) -> Union[torch.Tensor, np.ndarray]:
+ ) -> torch.Tensor | np.ndarray:
"""
This method generates output by executing PyTorch runtime or the compiled ``qpc`` on ``Cloud AI 100`` Hardware cards.
``Mandatory`` Args:
@@ -5509,7 +5573,7 @@ def cloud_ai_100_feature_generate(
self,
processor,
inputs: torch.Tensor,
- device_ids: List[int] = None,
+ device_ids: list[int] = None,
) -> np.ndarray:
"""
Generates features with list of prompts using AI 100 runtime.
@@ -5547,7 +5611,7 @@ def cloud_ai_100_feature_generate(
transcriptions = processor.batch_decode(torch.tensor(predicted_ids))
return transcriptions
- def pytorch_feature_generate(self, processor, model, inputs: Union[torch.Tensor, np.ndarray]) -> List[torch.Tensor]:
+ def pytorch_feature_generate(self, processor, model, inputs: torch.Tensor | np.ndarray) -> list[torch.Tensor]:
"""
Generates features from a list of text prompts using a PyTorch model.
diff --git a/QEfficient/transformers/models/qwen2_5_vl/modeling_qwen2_5_vl.py b/QEfficient/transformers/models/qwen2_5_vl/modeling_qwen2_5_vl.py
index 2959a8d0de..39aba65937 100644
--- a/QEfficient/transformers/models/qwen2_5_vl/modeling_qwen2_5_vl.py
+++ b/QEfficient/transformers/models/qwen2_5_vl/modeling_qwen2_5_vl.py
@@ -46,12 +46,67 @@
from QEfficient.utils.constants import MIN_MASKED_ATTENTION_VALUE
from QEfficient.utils.logging_utils import logger
-
-def qeff_prepare_mrope_cos_sin(cos, sin, position_ids):
+# def qeff_apply_interleaved_mrope(freqs, mrope_section):
+# """Apply interleaved MRoPE to 3D rotary embeddings.
+# Reorganizes frequency layout from chunked [TTT...HHH...WWW] to
+# interleaved [THWTHWTHW...TT], preserving frequency continuity.
+# args:
+# x: (3, bs, seq_len, head_dim // 2)
+# mrope_section: (3,)
+# returns:
+# x_t: (bs, seq_len, head_dim // 2)
+# """
+# freq_idx = torch.arange(freqs.shape[-1], device=freqs.device)
+# half_shape = freqs.shape[-1] // 2
+
+# h_mask = (freq_idx >= 1) & (freq_idx < mrope_section[1] * 3) & ((freq_idx - 1) % 3 == 0)
+# h_mask = h_mask | (
+# (freq_idx >= half_shape + 1)
+# & (freq_idx < half_shape + mrope_section[1] * 3)
+# & ((freq_idx - half_shape - 1) % 3 == 0)
+# )
+# w_mask = (freq_idx >= 2) & (freq_idx < mrope_section[2] * 3) & ((freq_idx - 2) % 3 == 0)
+# w_mask = w_mask | (
+# (freq_idx >= half_shape + 2)
+# & (freq_idx < half_shape + mrope_section[2] * 3)
+# & ((freq_idx - half_shape - 2) % 3 == 0)
+# )
+
+# freqs_t = torch.where(h_mask, freqs[1], freqs[0])
+# freqs_t = torch.where(w_mask, freqs[2], freqs_t)
+# return freqs_t
+
+
+def qeff_apply_interleaved_mrope(freqs, mrope_section):
+ """Apply interleaved MRoPE to 3D rotary embeddings.
+ Reorganizes frequency layout from chunked [TTT...HHH...WWW] to
+ interleaved [THWTHWTHW...TT], preserving frequency continuity.
+ args:
+ x: (3, bs, seq_len, head_dim // 2)
+ mrope_section: (3,)
+ returns:
+ x_t: (bs, seq_len, head_dim // 2)
+ """
+ freqs_t = freqs[0] # just overwrite the first dimension T
+ half_shape = freqs.shape[-1] // 2
+ for dim, offset in enumerate((1, 2), start=1): # H, W
+ length = mrope_section[dim] * 3
+ idx = slice(offset, length, 3)
+ freqs_t[..., idx] = freqs[dim, ..., idx]
+ offset += half_shape
+ length += half_shape
+ idx = slice(offset, length, 3)
+ freqs_t[..., idx] = freqs[dim, ..., idx]
+ return freqs_t
+
+
+def qeff_prepare_mrope_cos_sin(cos, sin, position_ids, mrope_section):
+ cos = cos.to(device=position_ids.device)
+ sin = sin.to(device=position_ids.device)
cos = cos[position_ids]
sin = sin[position_ids]
- cos = torch.cat([cos[0, ..., 0:32], cos[1, ..., 32:80], cos[2, ..., 80:128]], dim=-1).unsqueeze(1)
- sin = torch.cat([sin[0, ..., 0:32], sin[1, ..., 32:80], sin[2, ..., 80:128]], dim=-1).unsqueeze(1)
+ cos = qeff_apply_interleaved_mrope(cos, mrope_section).unsqueeze(1)
+ sin = qeff_apply_interleaved_mrope(sin, mrope_section).unsqueeze(1)
return cos, sin
@@ -95,6 +150,12 @@ def qeff_apply_rotary_pos_emb_vision(
return q_embed, k_embed
+def qeff_cumsum_dim1(tensor: torch.Tensor) -> torch.Tensor:
+ seq_len = tensor.shape[1]
+ cumsum_mask = torch.tril(torch.ones((seq_len, seq_len), dtype=tensor.dtype, device=tensor.device))
+ return tensor @ cumsum_mask
+
+
class QEffQwen2_5_VLAttentionMask(nn.Module):
"""Builds the windowed attention mask used by the vision blocks.
@@ -108,23 +169,17 @@ def forward(self, hidden_states: torch.Tensor, cu_seqlens: torch.Tensor) -> torc
dtype = hidden_states.dtype
min_val = torch.finfo(dtype).min
- # Create index grids
- rows = torch.arange(seq_len).view(1, -1)
- cols = torch.arange(seq_len).view(-1, 1)
+ rows = torch.arange(seq_len, device=hidden_states.device).view(1, -1)
+ cols = torch.arange(seq_len, device=hidden_states.device).view(-1, 1)
- # Prepare start and end indices
start = cu_seqlens[:-1].view(-1, 1, 1)
end = cu_seqlens[1:].view(-1, 1, 1)
- # Create block masks using broadcasting
row_mask = (rows >= start) & (rows < end)
col_mask = (cols >= start) & (cols < end)
- block_mask = row_mask & col_mask
-
- # Combine all blocks into one mask
- final_mask = torch.ones((seq_len, seq_len), dtype=dtype)
- final_mask[block_mask.any(dim=0)] = 0
- final_mask = torch.where(final_mask == 1.0, min_val, final_mask)
+ allowed_mask = (row_mask & col_mask).any(dim=0)
+ blocked_mask = torch.full((seq_len, seq_len), min_val, device=hidden_states.device, dtype=dtype)
+ final_mask = torch.where(allowed_mask, torch.zeros_like(blocked_mask), blocked_mask)
return final_mask.unsqueeze(0)
@@ -158,6 +213,8 @@ def forward(
sin = emb.sin()
else:
cos, sin = position_embeddings
+ cos = cos.to(device=q.device)
+ sin = sin.to(device=q.device)
q, k = qeff_apply_rotary_pos_emb_vision(q, k, cos, sin)
q = q.transpose(0, 1)
@@ -198,11 +255,11 @@ def __qeff_init__(self):
self.attention_mask_builder = QEffQwen2_5_VLAttentionMask()
def rot_pos_emb(self, grid_thw):
- pos_ids = []
-
bs, t, h, w = grid_thw.shape
+ inv_freq = self.rotary_pos_emb.inv_freq
+ device = inv_freq.device
- hpos_ids = torch.arange(h).unsqueeze(1).expand(-1, w)
+ hpos_ids = torch.arange(h, device=device).unsqueeze(1).expand(-1, w)
hpos_ids = hpos_ids.reshape(
h // self.spatial_merge_size,
self.spatial_merge_size,
@@ -212,7 +269,7 @@ def rot_pos_emb(self, grid_thw):
hpos_ids = hpos_ids.permute(0, 2, 1, 3)
hpos_ids = hpos_ids.flatten()
- wpos_ids = torch.arange(w).unsqueeze(0).expand(h, -1)
+ wpos_ids = torch.arange(w, device=device).unsqueeze(0).expand(h, -1)
wpos_ids = wpos_ids.reshape(
h // self.spatial_merge_size,
self.spatial_merge_size,
@@ -221,17 +278,10 @@ def rot_pos_emb(self, grid_thw):
)
wpos_ids = wpos_ids.permute(0, 2, 1, 3)
wpos_ids = wpos_ids.flatten()
- pos_ids.append(torch.stack([hpos_ids, wpos_ids], dim=-1).repeat(t, 1))
- pos_ids = torch.cat(pos_ids, dim=0)
-
- x_expanded = pos_ids.unsqueeze(0)
- x_expanded = x_expanded.expand(bs, -1, -1)
- pos_ids = x_expanded.reshape(-1, pos_ids.size(1))
-
- max_grid_size = max(grid_thw.shape)
- rotary_pos_emb_full = self.rotary_pos_emb(max_grid_size)
- rotary_pos_emb = rotary_pos_emb_full[pos_ids].flatten(1)
- return rotary_pos_emb
+ pos_ids = torch.stack([hpos_ids, wpos_ids], dim=-1).repeat(t, 1)
+ pos_ids = pos_ids.repeat(bs, 1) # check size of t
+ rotary_pos_emb = pos_ids.to(dtype=inv_freq.dtype).unsqueeze(-1) * inv_freq.view(1, 1, -1)
+ return rotary_pos_emb.flatten(1)
def get_window_index(self, grid_thw):
window_index: list = []
@@ -244,7 +294,9 @@ def get_window_index(self, grid_thw):
grid_h // self.spatial_merge_size,
grid_w // self.spatial_merge_size,
)
- index = torch.arange(grid_t * llm_grid_h * llm_grid_w).reshape(grid_t, llm_grid_h, llm_grid_w)
+ index = torch.arange(grid_t * llm_grid_h * llm_grid_w, device=grid_thw.device).reshape(
+ grid_t, llm_grid_h, llm_grid_w
+ )
pad_h = vit_merger_window_size - llm_grid_h % vit_merger_window_size
pad_w = vit_merger_window_size - llm_grid_w % vit_merger_window_size
@@ -269,15 +321,13 @@ def get_window_index(self, grid_thw):
seqlens = (index_padded != -100).sum([2, 3]).reshape(-1)
- x_expanded = seqlens.unsqueeze(0)
- x_expanded = x_expanded.expand(bs, -1)
- seqlens = x_expanded.reshape(-1)
+ seqlens = seqlens.repeat(bs)
index_padded = index_padded.reshape(-1)
mask = (index_padded == -100).to(torch.int32)
- if torch.jit.is_tracing():
+ if torch.jit.is_tracing() or torch._dynamo.is_compiling():
order = torch.argsort(mask)
else:
order = torch.argsort(mask, stable=True)
@@ -286,7 +336,7 @@ def get_window_index(self, grid_thw):
index_new = index_new[: index.reshape(-1).size(0)]
step = grid_t * llm_grid_h * llm_grid_w
- batch_indices = torch.arange(bs)
+ batch_indices = torch.arange(bs, device=grid_thw.device)
batch_indices = batch_indices.view(-1, 1)
offsets = batch_indices * step
window_index_tmp = index_new.unsqueeze(0) + offsets
@@ -294,7 +344,9 @@ def get_window_index(self, grid_thw):
cu_seqlens_tmp = seqlens.cumsum(0) * self.spatial_merge_unit + cu_window_seqlens[-1]
- cu_window_seqlens = torch.cat([torch.tensor([0], dtype=cu_seqlens_tmp.dtype), cu_seqlens_tmp])
+ cu_window_seqlens = torch.cat(
+ [torch.zeros(1, dtype=cu_seqlens_tmp.dtype, device=cu_seqlens_tmp.device), cu_seqlens_tmp]
+ )
return window_index, cu_window_seqlens
@@ -309,6 +361,7 @@ def forward(self, hidden_states: torch.Tensor, grid_thw: torch.Tensor) -> torch.
Returns:
`torch.Tensor`: hidden_states.
"""
+ grid_thw = grid_thw.to(device=hidden_states.device)
hidden_states = self.patch_embed(hidden_states)
rotary_pos_emb = self.rot_pos_emb(grid_thw)
@@ -316,7 +369,8 @@ def forward(self, hidden_states: torch.Tensor, grid_thw: torch.Tensor) -> torch.
window_index, cu_window_seqlens = self.get_window_index(grid_thw)
cu_window_seqlens = cu_window_seqlens.to(
- device=hidden_states.device, dtype=grid_thw.dtype if torch.jit.is_tracing() else torch.int32
+ device=hidden_states.device,
+ dtype=grid_thw.dtype if torch.jit.is_tracing() or torch._dynamo.is_compiling() else torch.int32,
)
# cu_window_seqlens = torch.unique_consecutive(cu_window_seqlens)
@@ -335,9 +389,9 @@ def forward(self, hidden_states: torch.Tensor, grid_thw: torch.Tensor) -> torch.
bs, t, h, w = grid_thw.shape
- t = torch.arange(t, t + 1).squeeze().expand(bs)
- h = torch.arange(h, h + 1).squeeze().expand(bs)
- w = torch.arange(w, w + 1).squeeze().expand(bs)
+ t = torch.arange(t, t + 1, device=grid_thw.device).squeeze().expand(bs)
+ h = torch.arange(h, h + 1, device=grid_thw.device).squeeze().expand(bs)
+ w = torch.arange(w, w + 1, device=grid_thw.device).squeeze().expand(bs)
cu_seqlens = (h * w).cumsum(
dim=0,
@@ -345,10 +399,10 @@ def forward(self, hidden_states: torch.Tensor, grid_thw: torch.Tensor) -> torch.
# - FA2 requires that cu_seqlens_q must have dtype int32
# - torch.onnx.export requires that cu_seqlens_q must have same dtype as grid_thw
# See https://github.com/huggingface/transformers/pull/34852 for more information
- dtype=grid_thw.dtype if torch.jit.is_tracing() else torch.int32,
+ dtype=grid_thw.dtype if torch.jit.is_tracing() or torch._dynamo.is_compiling() else torch.int32,
)
- cu_seqlens = torch.cat([torch.tensor([0], dtype=cu_seqlens.dtype), cu_seqlens])
+ cu_seqlens = torch.cat([torch.zeros(1, dtype=cu_seqlens.dtype, device=cu_seqlens.device), cu_seqlens])
full_attention_mask = self.attention_mask_builder(hidden_states, cu_seqlens)
window_attention_mask = self.attention_mask_builder(hidden_states, cu_window_seqlens)
@@ -382,9 +436,7 @@ class QEffQwen2_5_VLRotaryEmbedding(Qwen2_5_VLRotaryEmbedding):
def __init__(self, config: Qwen2_5_VLConfig, device=None):
super().__init__(config=config)
# Build here to make `torch.jit.trace` work.
- self._set_cos_sin_cache(
- seq_len=self.original_max_seq_len, device=self.inv_freq.device, dtype=torch.get_default_dtype()
- )
+ self._set_cos_sin_cache(seq_len=self.original_max_seq_len, device=self.inv_freq.device, dtype=config.dtype)
def _set_cos_sin_cache(self, seq_len, device, dtype):
self.max_seq_len_cached = seq_len
@@ -416,10 +468,11 @@ def eager_attention_forward(
attn_weights = torch.matmul(query, key_states.transpose(2, 3)) / math.sqrt(module.head_dim)
+ mask_value = torch.full_like(attn_weights, MIN_MASKED_ATTENTION_VALUE, dtype=attn_weights.dtype)
+
if attention_mask is not None:
- attn_weights = torch.where(
- attention_mask, torch.tensor(MIN_MASKED_ATTENTION_VALUE, dtype=module.config.torch_dtype), attn_weights
- )
+ # Apply the attention mask
+ attn_weights = torch.where(attention_mask, mask_value, attn_weights)
attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query.dtype)
attn_output = torch.matmul(attn_weights, value_states)
@@ -465,7 +518,6 @@ def forward(
query_states, key_states = qeff_apply_rotary_pos_emb(query_states, key_states, cos_cached, sin_cached)
- past_seen_tokens = past_key_values.get_seq_length() if past_key_values is not None else 0
blocking_config = getattr(self, "attn_blocking_config", AttentionBlockingConfig())
use_blocking = blocking_config is not None and (blocking_config.mode != BlockingMode.NONE)
if use_blocking:
@@ -586,8 +638,8 @@ def forward(
if output_attentions:
outputs += (self_attn_weights,)
- if use_cache:
- outputs += (present_key_value,)
+ # if use_cache:
+ # outputs += (present_key_value,)
return outputs
@@ -645,7 +697,9 @@ def forward(
# decoder layers
all_hidden_states = () if output_hidden_states else None
all_self_attns = () if output_attentions else None
- cos, sin = qeff_prepare_mrope_cos_sin(self.cos_cached, self.sin_cached, position_ids[1:])
+ cos, sin = qeff_prepare_mrope_cos_sin(
+ self.cos_cached, self.sin_cached, position_ids[1:], self.config.rope_scaling["mrope_section"]
+ )
for decoder_layer in self.layers:
if output_hidden_states:
@@ -760,8 +814,8 @@ def get_submodules_for_export(self) -> Type[nn.Module]:
def forward(self, pixel_values, image_grid_thw):
image_embeds = self.model.visual(pixel_values, grid_thw=image_grid_thw)
bs = image_grid_thw.shape[0]
- split_size = torch.floor_divide(torch.tensor(image_embeds.size(0)), bs)
- image_embeds = image_embeds.reshape(bs, split_size, image_embeds.size(1))
+ split_size = image_embeds.shape[0] // bs
+ image_embeds = image_embeds.reshape(bs, split_size, image_embeds.shape[-1])
return image_embeds
@@ -794,12 +848,15 @@ def forward(
inputs_embeds = self.model.get_input_embeddings()(input_ids)
B, N, C = inputs_embeds.shape
selected = input_ids == self.model.config.image_token_id
+ # indices1 = qeff_cumsum_dim1(selected.to(torch.int64)) - 1
indices1 = selected.to(torch.int64).cumsum(1) - 1
indices1 = torch.where(indices1 != -1, indices1 + image_idx, indices1)
- indices0 = torch.arange(selected.unsqueeze(0).shape[0]).view(-1, 1)
+ indices0 = torch.arange(selected.unsqueeze(0).shape[0], device=selected.device).view(-1, 1)
image_features_expanded = vision_embeds.reshape(-1, C).unsqueeze(0)[indices0, indices1]
image_input_embeds = torch.where(selected.unsqueeze(-1), image_features_expanded, inputs_embeds)
- inputs_embeds = torch.where(input_ids.shape[1] == torch.tensor(1), inputs_embeds, image_input_embeds)
+ inputs_embeds = torch.where(
+ input_ids.shape[1] == torch.tensor(1, device=input_ids.device), inputs_embeds, image_input_embeds
+ )
outputs = self.model.model(
inputs_embeds=inputs_embeds,
position_ids=position_ids,
@@ -810,12 +867,14 @@ def forward(
)
logit_index = position_ids[0].to(torch.int32).argmax(1, keepdim=True)
- hidden_states = outputs.last_hidden_state[torch.arange(position_ids[0].shape[0]).view(-1, 1), logit_index]
+ hidden_states = outputs.last_hidden_state[
+ torch.arange(position_ids[0].shape[0], device=position_ids.device).view(-1, 1), logit_index
+ ]
logits = self.model.lm_head(hidden_states)
logits = logits.float()
image_idx = (indices1.max() + 1).unsqueeze(0).unsqueeze(0)
- return logits, vision_embeds, image_idx, outputs.past_key_values
+ return logits, vision_embeds.clone(), image_idx, outputs.past_key_values
class QEffQwen_2_5_vl_ForConditionalGeneration(Qwen2_5_VLForConditionalGeneration):
@@ -832,49 +891,45 @@ def get_dummy_inputs(
continuous_batching: bool = False,
**kwargs,
):
+ bs: int = constants.ONNX_EXPORT_EXAMPLE_BATCH_SIZE + 1
+ fbs: int = constants.ONNX_EXPORT_EXAMPLE_FBS
+
prefill_seq_len = kwargs.get("prefill_seq_len", constants.ONNX_EXPORT_EXAMPLE_SEQ_LEN)
if prefill_seq_len is None:
prefill_seq_len = constants.ONNX_EXPORT_EXAMPLE_SEQ_LEN
prefill_seq_len = int(prefill_seq_len)
inputs_shapes = {}
- inputs_shapes["input_ids"] = (constants.ONNX_EXPORT_EXAMPLE_BATCH_SIZE, prefill_seq_len)
+ inputs_shapes["input_ids"] = (bs, prefill_seq_len)
vision_size = 3577
inputs_shapes["vision_embeds"] = (
- constants.ONNX_EXPORT_EXAMPLE_BATCH_SIZE,
+ bs,
vision_size,
self.model.config.text_config.hidden_size,
)
- inputs_shapes["image_grid_thw"] = (1, 1, 98, 146)
+ inputs_shapes["image_grid_thw"] = (bs, 1, 98, 146)
inputs_shapes["position_ids"] = (
3,
- constants.ONNX_EXPORT_EXAMPLE_BATCH_SIZE,
+ bs,
prefill_seq_len,
)
- inputs_shapes["pixel_values"] = (14308, 1176)
+ inputs_shapes["pixel_values"] = (14308 * bs, 1176)
inputs_shapes["image_idx"] = (1, 1)
- inputs_shapes["image_sizes"] = (constants.ONNX_EXPORT_EXAMPLE_BATCH_SIZE, 2)
+ inputs_shapes["image_sizes"] = (bs, 2)
# Define inputs
vision_inputs = {}
lang_inputs = {}
- vision_inputs["pixel_values"] = torch.zeros((inputs_shapes["pixel_values"]), dtype=self.config.torch_dtype)
+ vision_inputs["pixel_values"] = torch.zeros((inputs_shapes["pixel_values"]), dtype=self.config.dtype)
vision_inputs["image_grid_thw"] = torch.zeros((inputs_shapes["image_grid_thw"]), dtype=torch.int64)
lang_inputs["input_ids"] = torch.zeros((inputs_shapes["input_ids"]), dtype=torch.int64)
- lang_inputs["vision_embeds"] = torch.zeros((inputs_shapes["vision_embeds"]), dtype=self.config.torch_dtype)
+ lang_inputs["vision_embeds"] = torch.zeros((inputs_shapes["vision_embeds"]), dtype=self.config.dtype)
lang_inputs["position_ids"] = (
- (
- torch.arange(prefill_seq_len, dtype=torch.int64)
- .view(1, prefill_seq_len)
- .repeat(constants.ONNX_EXPORT_EXAMPLE_BATCH_SIZE, 1)
- )
+ (torch.arange(prefill_seq_len, dtype=torch.int64).view(1, prefill_seq_len).repeat(bs, 1))
.unsqueeze(0)
.repeat(4, 1, 1)
)
lang_inputs["image_idx"] = torch.zeros((inputs_shapes["image_idx"]), dtype=torch.int64)
- bs: int = constants.ONNX_EXPORT_EXAMPLE_BATCH_SIZE
- fbs: int = constants.ONNX_EXPORT_EXAMPLE_FBS
-
# Add data for KV
kv_cache_shape = get_padding_shape_from_config(
config=self.model.config.text_config,
@@ -885,7 +940,7 @@ def get_dummy_inputs(
lang_inputs["past_key_values"] = [[] for _ in range(self.model.config.text_config.num_hidden_layers)]
for i in range(self.model.config.text_config.num_hidden_layers):
for kv in ["key", "value"]:
- lang_inputs["past_key_values"][i].append(torch.zeros(kv_cache_shape, dtype=self.config.torch_dtype))
+ lang_inputs["past_key_values"][i].append(torch.zeros(kv_cache_shape, dtype=self.config.dtype))
if continuous_batching:
lang_inputs["batch_index"] = torch.arange(bs).view(bs, 1)
@@ -1072,19 +1127,24 @@ def get_specializations(
return lang, compiler_options
def get_onnx_dynamic_axes(
- self, comp_ctx_lengths: Optional[List[int]] = None, kv_offload: bool = False, continuous_batching: bool = False
+ self,
+ comp_ctx_lengths: Optional[List[int]] = None,
+ kv_offload: bool = False,
+ continuous_batching: bool = False,
+ batch_fold: bool = False,
):
# Define dynamic axes
num_layers = self.config.text_config.num_hidden_layers
+ batch_axis = "full_batch_size" if continuous_batching and batch_fold else "batch_size"
vision_dynamic_axes = {
- "pixel_values": {0: "grid_height", 1: "grid_width"},
+ "pixel_values": {0: "grid_height"},
"image_grid_thw": {0: "batch_size", 2: "grid_h", 3: "grid_w"},
}
lang_dynamic_axes = {
- "input_ids": {0: "batch_size", 1: "seq_len"},
- "position_ids": {1: "batch_size", 2: "seq_len"},
+ "input_ids": {0: batch_axis, 1: "seq_len"},
+ "position_ids": {1: batch_axis, 2: "seq_len"},
"vision_embeds": {0: "vision_batch_size", 1: "vision_size"},
}
@@ -1099,7 +1159,7 @@ def get_onnx_dynamic_axes(
}
if continuous_batching:
- lang_dynamic_axes["batch_index"] = {0: "batch_size"}
+ lang_dynamic_axes["batch_index"] = {0: batch_axis}
if comp_ctx_lengths is not None:
lang_dynamic_axes["comp_ctx_lengths"] = {0: "comp_ctx_lengths"}
@@ -1136,7 +1196,9 @@ def get_output_names(self, kv_offload: bool = False):
def prepare_inputs_for_generation(self, inputs, prefill_seq_len=128, batch_size=1):
input_ids_length = inputs["input_ids"].shape[1]
- inputs["position_ids"] = torch.arange(input_ids_length).view(1, 1, input_ids_length).expand(-1, batch_size, -1)
+ inputs["position_ids"] = torch.arange(input_ids_length, device=inputs["input_ids"].device).view(
+ 1, 1, input_ids_length
+ ).expand(-1, batch_size, -1)
mm_token_type_ids = inputs.get("mm_token_type_ids")
if mm_token_type_ids is None:
@@ -1174,7 +1236,7 @@ def get_inputs_info(self):
IOInfo(name="attention_mask", datatype=torch.int64, shape=("batch_size", "seq_len")),
IOInfo(
name="pixel_values",
- datatype=self.config.torch_dtype,
+ datatype=self.config.dtype,
shape=("batch_size", 3, "image_size", "image_size"),
),
]
diff --git a/QEfficient/transformers/models/qwen3_5/modeling_qwen3_5.py b/QEfficient/transformers/models/qwen3_5/modeling_qwen3_5.py
index 2cdcf781e2..41c195a44b 100644
--- a/QEfficient/transformers/models/qwen3_5/modeling_qwen3_5.py
+++ b/QEfficient/transformers/models/qwen3_5/modeling_qwen3_5.py
@@ -39,12 +39,17 @@
generic_blocked_attention_interface,
)
from QEfficient.customop import (
- CtxGatherFuncCB,
- CtxGatherFuncCB3D,
- CtxScatterFuncCB,
- CtxScatterFuncCB3D,
+ # CtxGatherFuncCB,
+ # CtxGatherFuncCB3D,
+ # CtxScatterFuncCB,
+ # CtxScatterFuncCB3D,
+ ctx_gather_cb,
+ ctx_gather_cb_3d,
+ ctx_scatter_cb,
+ ctx_scatter_cb_3d,
)
from QEfficient.customop.rms_norm import CustomRMSNormFunc
+from QEfficient.customop.utils import select_interface
from QEfficient.transformers.cache_utils import (
QEffDynamicLayer,
)
@@ -63,8 +68,9 @@ class QEffQwen3_5GatedDeltaNetCustomRMSNormAIC(nn.Module):
"""
def forward(self, hidden_states, gate):
+ rms_interface = select_interface(CustomRMSNormFunc.apply, torch.ops.qefficient.rms_norm)
return (
- CustomRMSNormFunc.apply(
+ rms_interface(
hidden_states, self.weight, self.variance_epsilon if hasattr(self, "variance_epsilon") else self.eps
)
) * F.silu(gate.to(torch.float32))
@@ -270,6 +276,8 @@ def qeff_apply_interleaved_mrope(freqs, mrope_section):
def qeff_prepare_mrope_cos_sin(cos, sin, position_ids, mrope_section, dtype=None):
+ cos = cos.to(device=position_ids.device)
+ sin = sin.to(device=position_ids.device)
invalid_pos_mask = position_ids < 0
safe_position_ids = torch.where(invalid_pos_mask, torch.zeros_like(position_ids), position_ids)
flat_pos = safe_position_ids.reshape(-1)
@@ -353,10 +361,15 @@ def eager_attention_forward(
value_states = repeat_kv(value, module.num_key_value_groups)
attn_weights = torch.matmul(query, key_states.transpose(2, 3)) * scaling
+ # if attention_mask is not None:
+ # attn_weights = torch.where(
+ # attention_mask, torch.tensor(MIN_MASKED_ATTENTION_VALUE, dtype=torch.float32), attn_weights
+ # )
+ mask_value = torch.full_like(attn_weights, MIN_MASKED_ATTENTION_VALUE, dtype=attn_weights.dtype)
+
if attention_mask is not None:
- attn_weights = torch.where(
- attention_mask, torch.tensor(MIN_MASKED_ATTENTION_VALUE, dtype=torch.float32), attn_weights
- )
+ # Apply the attention mask
+ attn_weights = torch.where(attention_mask, mask_value, attn_weights)
attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query.dtype)
attn_output = torch.matmul(attn_weights, value_states)
@@ -544,6 +557,9 @@ def torch_chunk_gated_delta_rule_qeff(
ones_lower=None,
eye=None,
):
+
+ if ones_lower is not None:
+ ones_lower = ones_lower.clone()
initial_dtype = query.dtype
# if use_qk_l2norm_in_kernel:
# query = l2norm(query, dim=-1, eps=1e-6)
@@ -588,15 +604,15 @@ def torch_chunk_gated_delta_rule_qeff(
scale = 1 / (query.shape[-1] ** 0.5)
query = query * scale
- v_beta = value * beta.unsqueeze(-1)
- k_beta = key * beta.unsqueeze(-1)
+ v_beta = value * beta.unsqueeze(-1).to(g.device)
+ k_beta = key * beta.unsqueeze(-1).to(g.device)
# reshape to chunks
query, key, value, k_beta, v_beta = [
x.reshape(x.shape[0], x.shape[1], -1, chunk_size, x.shape[-1]) for x in (query, key, value, k_beta, v_beta)
]
g = g.reshape(g.shape[0], g.shape[1], -1, chunk_size)
# mask = torch.triu(torch.ones(chunk_size, chunk_size, dtype=torch.bool, device=query.device), diagonal=0)
- mask = mask_causal
+ mask = mask_causal.to(g.device)
#
# chunk decay
@@ -612,10 +628,10 @@ def torch_chunk_gated_delta_rule_qeff(
# decay_mask = ((g.unsqueeze(-1) - g.unsqueeze(-2)).tril().exp().float()).tril() # original decay_mask
diff = g.unsqueeze(-1) - g.unsqueeze(-2) # (B, H, num_chunks, C, C)
+ mask_strict = mask_strict.to(diff.device)
diff = diff * (~mask_strict).float() # zero upper triangle (strict)
decay_mask = diff.exp().float()
- decay_mask = decay_mask * (~mask_strict).float() # ensure upper is zero
-
+ decay_mask = decay_mask * (~mask_strict).float().to(g.device) # ensure upper is zero
attn = -((k_beta @ key.transpose(-1, -2)) * decay_mask).masked_fill(mask, 0)
for i in range(1, chunk_size):
row = attn[..., i, :i].clone()
@@ -666,7 +682,7 @@ def torch_chunk_gated_delta_rule_qeff(
k_cumdecay = attn @ (k_beta * g.exp().unsqueeze(-1))
last_recurrent_state = (
- torch.zeros(batch_size, num_heads, k_head_dim, v_head_dim).to(value)
+ torch.zeros(batch_size, num_heads, k_head_dim, v_head_dim, device=value.device, dtype=value.dtype)
if initial_state is None
else initial_state.to(value)
)
@@ -694,6 +710,7 @@ def torch_chunk_gated_delta_rule_qeff(
)
core_attn_out = core_attn_out[:, :, :sequence_length]
core_attn_out = core_attn_out.transpose(1, 2).contiguous().to(initial_dtype)
+ core_attn_out = core_attn_out.clone()
return core_attn_out, last_recurrent_state
def _recurrent_step_batched(self, query, key, value, g, beta, recurrent_state):
@@ -767,13 +784,13 @@ def forward(
conv_ctx_indices = torch.arange(
conv_state_all.shape[1], dtype=torch.int64, device=conv_state_all.device
)[None, :]
- conv_state = CtxGatherFuncCB3D.apply(conv_state_all, conv_batch_index, conv_ctx_indices)
+ conv_state = ctx_gather_cb_3d(conv_state_all, conv_batch_index, conv_ctx_indices)
recurrent_batch_index = batch_index.to(recurrent_state_all.device)
recurrent_ctx_indices = torch.arange(
recurrent_state_all.shape[2], dtype=torch.int64, device=recurrent_state_all.device
)[None, None, :]
- recurrent_state = CtxGatherFuncCB.apply(
+ recurrent_state = ctx_gather_cb(
recurrent_state_all, recurrent_batch_index, recurrent_ctx_indices, recurrent_state_all.shape[2]
)
else:
@@ -792,7 +809,7 @@ def forward(
conv_position_ids = torch.arange(
conv_state_all.shape[1], dtype=torch.int64, device=conv_state_all.device
)[None, :]
- cache_params.conv_states[self.layer_idx] = CtxScatterFuncCB3D.apply(
+ cache_params.conv_states[self.layer_idx] = ctx_scatter_cb_3d(
conv_state_all, conv_batch_index, conv_position_ids, new_conv_state
)
else:
@@ -842,7 +859,9 @@ def forward(
# Select based on seq_len
# is_decode is SCALAR — torch.where broadcasts efficiently
# HW predicates entire branch at runtime
- is_decode = hidden_states.shape[1] == torch.tensor(1)
+ position_scalar = position_ids.reshape(-1)[0]
+ seq_len_tensor = torch.full_like(position_scalar, hidden_states.shape[1])
+ is_decode = seq_len_tensor == torch.ones_like(seq_len_tensor)
core_attn_out = torch.where(is_decode, recurrent_out, chunk_out)
last_recurrent_state = torch.where(is_decode, recurrent_S, chunk_S)
@@ -852,7 +871,7 @@ def forward(
recurrent_position_ids = torch.arange(
recurrent_state_all.shape[2], dtype=torch.int64, device=recurrent_state_all.device
)[None, :].expand(recurrent_batch_index.shape[0], -1)
- cache_params.recurrent_states[self.layer_idx] = CtxScatterFuncCB.apply(
+ cache_params.recurrent_states[self.layer_idx] = ctx_scatter_cb(
recurrent_state_all,
recurrent_batch_index,
recurrent_position_ids,
@@ -1153,7 +1172,10 @@ def forward(
else:
text_position_ids = position_ids[0] if position_ids.ndim == 3 else position_ids
logit_index = text_position_ids.to(torch.int32).argmax(1, keepdim=True)
- hidden_states = outputs.last_hidden_state[torch.arange(text_position_ids.shape[0]).view(-1, 1), logit_index]
+ hidden_states = outputs.last_hidden_state[
+ torch.arange(text_position_ids.shape[0], device=outputs.last_hidden_state.device).view(-1, 1),
+ logit_index,
+ ]
logits = self.lm_head(hidden_states).float()
return CausalLMOutputWithPast(
@@ -1242,10 +1264,6 @@ def rot_pos_emb(self, grid_thw: torch.Tensor) -> torch.Tensor:
freq_table = self.rotary_pos_emb(max_hw)
device = freq_table.device
bs, num_frames, height, width = grid_thw.shape
- grid_thw = (torch.tensor(grid_thw.shape, dtype=torch.int64)).unsqueeze(0)
-
- total_tokens = int(torch.prod(grid_thw, dim=1).sum().item())
- pos_ids = torch.empty((total_tokens, 2), dtype=torch.long, device=device)
merged_h, merged_w = height // merge_size, width // merge_size
@@ -1265,22 +1283,26 @@ def rot_pos_emb(self, grid_thw: torch.Tensor) -> torch.Tensor:
if num_frames > 1:
coords = coords.repeat(num_frames, 1)
- pos_ids = coords
- embeddings = freq_table[pos_ids]
+ coords = coords.repeat(bs, 1)
+ embeddings = freq_table[coords]
embeddings = embeddings.flatten(1)
return embeddings
def fast_pos_embed_interpolate(self, grid_thw):
bs, t, h, w = grid_thw.shape
- h_idxs = torch.linspace(0, self.num_grid_per_side - 1, h)
- w_idxs = torch.linspace(0, self.num_grid_per_side - 1, w)
+ device = self.pos_embed.weight.device
+ h_den = torch.clamp(torch.scalar_tensor(h - 1, device=device, dtype=torch.float32), min=1.0)
+ w_den = torch.clamp(torch.scalar_tensor(w - 1, device=device, dtype=torch.float32), min=1.0)
+ h_idxs = torch.arange(h, device=device, dtype=torch.float32) * ((self.num_grid_per_side - 1) / h_den)
+ w_idxs = torch.arange(w, device=device, dtype=torch.float32) * ((self.num_grid_per_side - 1) / w_den)
h_idxs_floor = h_idxs.int()
w_idxs_floor = w_idxs.int()
- max_t = torch.tensor(self.num_grid_per_side - 1, device=h_idxs.device)
- h_idxs_ceil = torch.minimum(h_idxs_floor + 1, max_t)
- w_idxs_ceil = torch.minimum(w_idxs_floor + 1, max_t)
+ max_idx_h = torch.full_like(h_idxs_floor, self.num_grid_per_side - 1)
+ max_idx_w = torch.full_like(w_idxs_floor, self.num_grid_per_side - 1)
+ h_idxs_ceil = torch.minimum(h_idxs_floor + 1, max_idx_h)
+ w_idxs_ceil = torch.minimum(w_idxs_floor + 1, max_idx_w)
dh = h_idxs - h_idxs_floor
dw = w_idxs - w_idxs_floor
@@ -1302,7 +1324,7 @@ def fast_pos_embed_interpolate(self, grid_thw):
(dh[None].T * dw[None]).flatten(),
]
- idx_tensor = torch.stack(indices, dim=0).to(dtype=torch.long, device=self.pos_embed.weight.device)
+ idx_tensor = torch.stack(indices, dim=0).to(dtype=torch.long, device=self.pos_embed.weight.device) # [4, h*w]
weight_tensor = torch.stack(weights, dim=0).to(
dtype=self.pos_embed.weight.dtype, device=self.pos_embed.weight.device
@@ -1310,11 +1332,8 @@ def fast_pos_embed_interpolate(self, grid_thw):
pos_embeds = self.pos_embed(idx_tensor) * weight_tensor[:, :, None]
patch_pos_embeds = pos_embeds[0] + pos_embeds[1] + pos_embeds[2] + pos_embeds[3]
- patch_pos_embeds = patch_pos_embeds.split([h * w])
-
- patch_pos_embeds_permute = []
merge_size = self.config.spatial_merge_size
- pos_embed = patch_pos_embeds[0]
+ pos_embed = patch_pos_embeds
pos_embed = pos_embed.repeat(t, 1)
pos_embed = (
@@ -1322,8 +1341,7 @@ def fast_pos_embed_interpolate(self, grid_thw):
.permute(0, 1, 3, 2, 4, 5)
.flatten(0, 4)
)
- patch_pos_embeds_permute.append(pos_embed)
- patch_pos_embeds = torch.cat(patch_pos_embeds_permute)
+ patch_pos_embeds = pos_embed
x_expanded = patch_pos_embeds.unsqueeze(0)
x_expanded = x_expanded.expand(bs, -1, -1)
patch_pos_embeds = x_expanded.reshape(-1, patch_pos_embeds.size(1))
@@ -1343,15 +1361,15 @@ def forward(self, hidden_states: torch.Tensor, grid_thw: torch.Tensor) -> torch.
position_embeddings = (emb.cos(), emb.sin())
bs, t, h, w = grid_thw.shape
- t = torch.arange(t, t + 1).squeeze().expand(bs)
- h = torch.arange(h, h + 1).squeeze().expand(bs)
- w = torch.arange(w, w + 1).squeeze().expand(bs)
+ t = torch.arange(t, t + 1, device=grid_thw.device).squeeze().expand(bs)
+ h = torch.arange(h, h + 1, device=grid_thw.device).squeeze().expand(bs)
+ w = torch.arange(w, w + 1, device=grid_thw.device).squeeze().expand(bs)
cu_seqlens = (h * w).cumsum(
dim=0,
dtype=torch.int32,
)
- cu_seqlens = torch.cat([torch.tensor([0], dtype=cu_seqlens.dtype), cu_seqlens])
+ cu_seqlens = torch.cat([torch.zeros(1, dtype=cu_seqlens.dtype, device=cu_seqlens.device), cu_seqlens])
for blk in self.blocks:
hidden_states = blk(
@@ -1392,25 +1410,32 @@ def forward(
sin = emb.sin()
else:
cos, sin = position_embeddings
+ cos = cos.to(device=q.device)
+ sin = sin.to(device=q.device)
q, k = apply_rotary_pos_emb_vision(q, k, cos, sin)
attention_mask = torch.full(
[1, seq_length, seq_length], torch.finfo(q.dtype).min, device=q.device, dtype=q.dtype
)
seq_len = attention_mask.shape[-1]
- rows = torch.arange(seq_len).view(1, -1)
- cols = torch.arange(seq_len).view(-1, 1)
+ rows = torch.arange(seq_len, device=q.device).view(1, -1)
+ cols = torch.arange(seq_len, device=q.device).view(-1, 1)
start = cu_seqlens[:-1].view(-1, 1, 1)
end = cu_seqlens[1:].view(-1, 1, 1)
row_mask = (rows >= start) & (rows < end)
col_mask = (cols >= start) & (cols < end)
- block_mask = row_mask & col_mask
+ # block_mask = row_mask & col_mask
- final_mask = torch.ones((seq_len, seq_len), dtype=torch.float32)
- final_mask[block_mask.any(dim=0)] = 0
- final_mask = torch.where(final_mask == 1.0, torch.finfo(q.dtype).min, final_mask)
- attention_mask[0] = final_mask
+ # final_mask = torch.ones((seq_len, seq_len), dtype=torch.float32)
+ # final_mask[block_mask.any(dim=0)] = 0
+ # final_mask = torch.where(final_mask == 1.0, torch.finfo(q.dtype).min, final_mask)
+ # attention_mask[0] = final_mask
+
+ allowed = (row_mask & col_mask).any(dim=0)
+ blocked_mask = torch.full((seq_len, seq_len), torch.finfo(q.dtype).min, device=q.device, dtype=q.dtype)
+ final_mask = torch.where(allowed, torch.zeros_like(blocked_mask), blocked_mask)
+ attention_mask = final_mask.unsqueeze(0)
q = q.transpose(0, 1)
k = k.transpose(0, 1)
@@ -1449,7 +1474,8 @@ def forward(self, pixel_values, image_grid_thw):
image_embeds = image_outputs.pooler_output
image_embeds = torch.cat(image_embeds, dim=0).to(pixel_values.device, pixel_values.dtype)
bs = image_grid_thw.shape[0]
- split_size = torch.floor_divide(torch.tensor(image_embeds.size(0)), bs)
+ # split_size = torch.floor_divide(torch.tensor(image_embeds.shape[0]), bs)
+ split_size = image_embeds.shape[0] // bs
image_embeds = image_embeds.reshape(bs, split_size, image_embeds.size(1))
return image_embeds
@@ -1479,7 +1505,7 @@ def forward(
selected = input_ids == self.model.config.image_token_id
indices1 = selected.to(torch.int64).cumsum(1) - 1
indices1 = torch.where(indices1 != -1, indices1 + image_idx, indices1)
- indices0 = torch.arange(selected.unsqueeze(0).shape[0]).view(-1, 1)
+ indices0 = torch.arange(selected.unsqueeze(0).shape[0], device=input_ids.device).view(-1, 1)
image_features_expanded = vision_embeds.reshape(-1, channel_size).unsqueeze(0)[indices0, indices1]
image_input_embeds = torch.where(selected.unsqueeze(-1), image_features_expanded, inputs_embeds)
inputs_embeds = image_input_embeds
@@ -1492,9 +1518,12 @@ def forward(
use_cache=True,
)
logit_index = position_ids[0].to(torch.int32).argmax(1, keepdim=True)
- hidden_states = outputs.last_hidden_state[torch.arange(position_ids[0].shape[0]).view(-1, 1), logit_index]
+ hidden_states = outputs.last_hidden_state[
+ torch.arange(position_ids[0].shape[0], device=outputs.last_hidden_state.device).view(-1, 1), logit_index
+ ]
logits = self.model.lm_head(hidden_states)
image_idx = (indices1.max() + 1).unsqueeze(0).unsqueeze(0)
+ vision_embeds = vision_embeds.clone()
return logits, vision_embeds, image_idx, outputs.past_key_values[: len(past_key_values)]
@@ -1593,7 +1622,9 @@ def forward(
hidden_states = outputs[0]
logit_index = position_ids[0].to(torch.int32).argmax(1, keepdim=True)
- hidden_states = outputs.last_hidden_state[torch.arange(position_ids[0].shape[0]).view(-1, 1), logit_index]
+ hidden_states = outputs.last_hidden_state[
+ torch.arange(position_ids[0].shape[0], device=outputs.last_hidden_state.device).view(-1, 1), logit_index
+ ]
logits = self.lm_head(hidden_states)
return logits, outputs.past_key_values[: len(past_key_values)]
@@ -1739,7 +1770,7 @@ def get_onnx_dynamic_axes(
vision_dynamic_axes = {
"pixel_values": {0: "grid_height", 1: "grid_width"},
- "image_grid_thw": {0: "batch_size", 1: "time", 2: "grid_h", 3: "grid_w"},
+ "image_grid_thw": {0: "batch_size", 2: "grid_h", 3: "grid_w"},
}
lang_dynamic_axes = {
@@ -1781,24 +1812,26 @@ def get_dummy_inputs(
**kwargs,
):
inputs_shapes = {}
-
- dummy_seq_len = 32
- inputs_shapes["input_ids"] = (constants.ONNX_EXPORT_EXAMPLE_BATCH_SIZE, dummy_seq_len)
+ bs: int = constants.ONNX_EXPORT_EXAMPLE_BATCH_SIZE + 1
+ fbs: int = constants.ONNX_EXPORT_EXAMPLE_FBS
+ dummy_seq_len = kwargs.get("prefill_seq_len", constants.ONNX_EXPORT_EXAMPLE_SEQ_LEN)
+ dummy_ctx_len = kwargs.get("ctx_len") or dummy_seq_len
+ inputs_shapes["input_ids"] = (bs, dummy_seq_len)
inputs_shapes["position_ids"] = (
4,
- constants.ONNX_EXPORT_EXAMPLE_BATCH_SIZE,
+ bs,
dummy_seq_len,
)
- inputs_shapes["pixel_values"] = (11008, 1536)
+ inputs_shapes["pixel_values"] = (11008 * bs, 1536)
inputs_shapes["image_grid_thw"] = (
- constants.ONNX_EXPORT_EXAMPLE_BATCH_SIZE,
+ bs,
1,
86,
128,
)
inputs_shapes["vision_embeds"] = (
- constants.ONNX_EXPORT_EXAMPLE_BATCH_SIZE,
+ bs,
2752,
self.model.config.text_config.hidden_size,
)
@@ -1811,23 +1844,16 @@ def get_dummy_inputs(
lang_inputs["input_ids"] = torch.zeros((inputs_shapes["input_ids"]), dtype=torch.int64)
lang_inputs["vision_embeds"] = torch.zeros((inputs_shapes["vision_embeds"]), dtype=torch.float32)
lang_inputs["position_ids"] = (
- (
- torch.arange(dummy_seq_len, dtype=torch.int64)
- .view(1, dummy_seq_len)
- .repeat(constants.ONNX_EXPORT_EXAMPLE_BATCH_SIZE, 1)
- )
+ (torch.arange(dummy_seq_len, dtype=torch.int64).view(1, dummy_seq_len).repeat(bs, 1))
.unsqueeze(0)
.repeat(4, 1, 1)
)
lang_inputs["image_idx"] = torch.zeros((inputs_shapes["image_idx"]), dtype=torch.int64)
- bs: int = constants.ONNX_EXPORT_EXAMPLE_BATCH_SIZE
- fbs: int = constants.ONNX_EXPORT_EXAMPLE_FBS
-
kv_cache_shape = get_padding_shape_from_config(
config=self.model.config.text_config,
batch_size=fbs if continuous_batching else bs,
- seq_len=dummy_seq_len,
+ seq_len=dummy_ctx_len,
)
linear_batch_size = fbs if continuous_batching else bs
@@ -1890,7 +1916,11 @@ def get_inputs_info(self):
def prepare_inputs_for_generation(self, inputs, prefill_seq_len=32, batch_size=1):
input_ids_length = inputs["input_ids"].shape[1]
- inputs["position_ids"] = torch.arange(input_ids_length).view(1, 1, input_ids_length).expand(-1, batch_size, -1)
+ inputs["position_ids"] = (
+ torch.arange(input_ids_length, device=inputs["input_ids"].device)
+ .view(1, 1, input_ids_length)
+ .expand(-1, batch_size, -1)
+ )
pos_ids, rope_deltas = self.model.get_rope_index(
inputs["input_ids"],
inputs["mm_token_type_ids"],
diff --git a/QEfficient/transformers/models/qwen3_vl/modeling_qwen3_vl.py b/QEfficient/transformers/models/qwen3_vl/modeling_qwen3_vl.py
index 79493bf226..5675957683 100644
--- a/QEfficient/transformers/models/qwen3_vl/modeling_qwen3_vl.py
+++ b/QEfficient/transformers/models/qwen3_vl/modeling_qwen3_vl.py
@@ -58,6 +58,37 @@ def _should_export_embedding_output(module) -> bool:
return False
+# def qeff_apply_interleaved_mrope(freqs, mrope_section):
+# """Apply interleaved MRoPE to 3D rotary embeddings.
+# Reorganizes frequency layout from chunked [TTT...HHH...WWW] to
+# interleaved [THWTHWTHW...TT], preserving frequency continuity.
+# args:
+# x: (3, bs, seq_len, head_dim // 2)
+# mrope_section: (3,)
+# returns:
+# x_t: (bs, seq_len, head_dim // 2)
+# """
+# freq_idx = torch.arange(freqs.shape[-1], device=freqs.device)
+# half_shape = freqs.shape[-1] // 2
+
+# h_mask = (freq_idx >= 1) & (freq_idx < mrope_section[1] * 3) & ((freq_idx - 1) % 3 == 0)
+# h_mask = h_mask | (
+# (freq_idx >= half_shape + 1)
+# & (freq_idx < half_shape + mrope_section[1] * 3)
+# & ((freq_idx - half_shape - 1) % 3 == 0)
+# )
+# w_mask = (freq_idx >= 2) & (freq_idx < mrope_section[2] * 3) & ((freq_idx - 2) % 3 == 0)
+# w_mask = w_mask | (
+# (freq_idx >= half_shape + 2)
+# & (freq_idx < half_shape + mrope_section[2] * 3)
+# & ((freq_idx - half_shape - 2) % 3 == 0)
+# )
+
+# freqs_t = torch.where(h_mask, freqs[1], freqs[0])
+# freqs_t = torch.where(w_mask, freqs[2], freqs_t)
+# return freqs_t
+
+
def qeff_apply_interleaved_mrope(freqs, mrope_section):
"""Apply interleaved MRoPE to 3D rotary embeddings.
Reorganizes frequency layout from chunked [TTT...HHH...WWW] to
@@ -82,6 +113,8 @@ def qeff_apply_interleaved_mrope(freqs, mrope_section):
def qeff_prepare_mrope_cos_sin(cos, sin, position_ids, mrope_section):
+ cos = cos.to(device=position_ids.device)
+ sin = sin.to(device=position_ids.device)
cos = cos[position_ids]
sin = sin[position_ids]
cos = qeff_apply_interleaved_mrope(cos, mrope_section).unsqueeze(1)
@@ -115,6 +148,12 @@ def qeff_apply_rotary_pos_emb(q, k, cos, sin):
return q_embed.to(q.dtype), k_embed.to(k.dtype)
+def qeff_cumsum_dim1(tensor: torch.Tensor) -> torch.Tensor:
+ seq_len = tensor.shape[1]
+ cumsum_mask = torch.tril(torch.ones((seq_len, seq_len), dtype=tensor.dtype, device=tensor.device))
+ return tensor @ cumsum_mask
+
+
class QEffQwen3VLTextRotaryEmbedding(Qwen3VLTextRotaryEmbedding):
"""
Copied from LlamaForCausalLM: https://github.com/huggingface/transformers/blob/main/src/transformers/models/llama/modeling_llama.py
@@ -125,9 +164,7 @@ class QEffQwen3VLTextRotaryEmbedding(Qwen3VLTextRotaryEmbedding):
def __init__(self, config: Qwen3VLTextConfig, device=None):
super().__init__(config=config)
# Build here to make `torch.jit.trace` work.
- self._set_cos_sin_cache(
- seq_len=self.original_max_seq_len, device=self.inv_freq.device, dtype=config.torch_dtype
- )
+ self._set_cos_sin_cache(seq_len=self.original_max_seq_len, device=self.inv_freq.device, dtype=config.dtype)
def _set_cos_sin_cache(self, seq_len, device, dtype):
self.max_seq_len_cached = seq_len
@@ -144,14 +181,18 @@ class QEffQwen3VLVisionModel(Qwen3VLVisionModel):
def rot_pos_emb(self, grid_thw: torch.Tensor) -> torch.Tensor:
merge_size = self.spatial_merge_size
- max_hw = max(grid_thw.shape)
- freq_table = self.rotary_pos_emb(max_hw) # (max_hw, dim // 2)
- device = freq_table.device
- bs, num_frames, height, width = grid_thw.shape
- grid_thw = (torch.tensor(grid_thw.shape, dtype=torch.int64)).unsqueeze(0)
+ # max_hw = max(grid_thw.shape)
+ # freq_table = self.rotary_pos_emb(max_hw) # (max_hw, dim // 2)
+ # device = freq_table.device
+ # bs, num_frames, height, width = grid_thw.shape
+ # grid_thw = (torch.tensor(grid_thw.shape, dtype=torch.int64)).unsqueeze(0)
- total_tokens = int(torch.prod(grid_thw, dim=1).sum().item())
- pos_ids = torch.empty((total_tokens, 2), dtype=torch.long, device=device)
+ # total_tokens = int(torch.prod(grid_thw, dim=1).sum().item())
+ # pos_ids = torch.empty((total_tokens, 2), dtype=torch.long, device=device)
+
+ bs, num_frames, height, width = grid_thw.shape
+ inv_freq = self.rotary_pos_emb.inv_freq
+ device = inv_freq.device
merged_h, merged_w = height // merge_size, width // merge_size
@@ -169,25 +210,27 @@ def rot_pos_emb(self, grid_thw: torch.Tensor) -> torch.Tensor:
coords = torch.stack((row_idx, col_idx), dim=-1)
- if num_frames > 1:
- coords = coords.repeat(num_frames, 1)
+ coords = coords.repeat(num_frames, 1)
- pos_ids = coords
- embeddings = freq_table[pos_ids] # lookup rotary embeddings
+ coords = coords.repeat(bs, 1)
+ embeddings = coords.to(dtype=inv_freq.dtype).unsqueeze(-1) * inv_freq.view(1, 1, -1)
embeddings = embeddings.flatten(1)
return embeddings
def fast_pos_embed_interpolate(self, grid_thw):
bs, t, h, w = grid_thw.shape
- h_idxs = torch.linspace(0, self.num_grid_per_side - 1, h)
- w_idxs = torch.linspace(0, self.num_grid_per_side - 1, w)
+ device = self.pos_embed.weight.device
+ h_den = torch.clamp(torch.scalar_tensor(h - 1, device=device, dtype=torch.float32), min=1.0)
+ w_den = torch.clamp(torch.scalar_tensor(w - 1, device=device, dtype=torch.float32), min=1.0)
+ h_idxs = torch.arange(h, device=device, dtype=torch.float32) * ((self.num_grid_per_side - 1) / h_den)
+ w_idxs = torch.arange(w, device=device, dtype=torch.float32) * ((self.num_grid_per_side - 1) / w_den)
h_idxs_floor = h_idxs.int()
w_idxs_floor = w_idxs.int()
- max_t = torch.tensor(self.num_grid_per_side - 1, device=h_idxs.device)
-
- h_idxs_ceil = torch.minimum(h_idxs_floor + 1, max_t) # working
- w_idxs_ceil = torch.minimum(w_idxs_floor + 1, max_t)
+ max_idx_h = torch.full_like(h_idxs_floor, self.num_grid_per_side - 1)
+ max_idx_w = torch.full_like(w_idxs_floor, self.num_grid_per_side - 1)
+ h_idxs_ceil = torch.minimum(h_idxs_floor + 1, max_idx_h)
+ w_idxs_ceil = torch.minimum(w_idxs_floor + 1, max_idx_w)
dh = h_idxs - h_idxs_floor
dw = w_idxs - w_idxs_floor
@@ -217,11 +260,8 @@ def fast_pos_embed_interpolate(self, grid_thw):
pos_embeds = self.pos_embed(idx_tensor) * weight_tensor[:, :, None]
patch_pos_embeds = pos_embeds[0] + pos_embeds[1] + pos_embeds[2] + pos_embeds[3]
- patch_pos_embeds = patch_pos_embeds.split([h * w])
-
- patch_pos_embeds_permute = []
merge_size = self.config.spatial_merge_size
- pos_embed = patch_pos_embeds[0]
+ pos_embed = patch_pos_embeds
pos_embed = pos_embed.repeat(t, 1)
pos_embed = (
@@ -229,14 +269,14 @@ def fast_pos_embed_interpolate(self, grid_thw):
.permute(0, 1, 3, 2, 4, 5)
.flatten(0, 4)
)
- patch_pos_embeds_permute.append(pos_embed)
- patch_pos_embeds = torch.cat(patch_pos_embeds_permute)
+ patch_pos_embeds = pos_embed
x_expanded = patch_pos_embeds.unsqueeze(0)
x_expanded = x_expanded.expand(bs, -1, -1)
patch_pos_embeds = x_expanded.reshape(-1, patch_pos_embeds.size(1))
return patch_pos_embeds
def forward(self, hidden_states: torch.Tensor, grid_thw: torch.Tensor) -> torch.Tensor:
+ grid_thw = grid_thw.to(device=hidden_states.device)
hidden_states = self.patch_embed(hidden_states)
pos_embeds = self.fast_pos_embed_interpolate(grid_thw)
@@ -251,15 +291,15 @@ def forward(self, hidden_states: torch.Tensor, grid_thw: torch.Tensor) -> torch.
position_embeddings = (emb.cos(), emb.sin())
bs, t, h, w = grid_thw.shape
- t = torch.arange(t, t + 1).squeeze().expand(bs)
- h = torch.arange(h, h + 1).squeeze().expand(bs)
- w = torch.arange(w, w + 1).squeeze().expand(bs)
+ t = torch.arange(t, t + 1, device=grid_thw.device).squeeze().expand(bs)
+ h = torch.arange(h, h + 1, device=grid_thw.device).squeeze().expand(bs)
+ w = torch.arange(w, w + 1, device=grid_thw.device).squeeze().expand(bs)
cu_seqlens = (h * w).cumsum(
dim=0,
dtype=torch.int32,
)
- cu_seqlens = torch.cat([torch.tensor([0], dtype=cu_seqlens.dtype), cu_seqlens])
+ cu_seqlens = torch.cat([torch.full_like(cu_seqlens[:1], 0), cu_seqlens])
deepstack_feature_lists = []
for layer_num, blk in enumerate(self.blocks):
@@ -306,33 +346,23 @@ def forward(
sin = emb.sin()
else:
cos, sin = position_embeddings
+ cos = cos.to(device=q.device)
+ sin = sin.to(device=q.device)
q, k = apply_rotary_pos_emb_vision(q, k, cos, sin)
- attention_mask = torch.full(
- [1, seq_length, seq_length], torch.finfo(q.dtype).min, device=q.device, dtype=q.dtype
- )
-
- # Create index grids
- seq_len = attention_mask.shape[-1]
- rows = torch.arange(seq_len).view(1, -1)
- cols = torch.arange(seq_len).view(-1, 1)
+ seq_len = seq_length
+ rows = torch.arange(seq_len, device=q.device).view(1, -1)
+ cols = torch.arange(seq_len, device=q.device).view(-1, 1)
- # Prepare start and end indices
start = cu_seqlens[:-1].view(-1, 1, 1)
end = cu_seqlens[1:].view(-1, 1, 1)
- # Create block masks using broadcasting
row_mask = (rows >= start) & (rows < end)
col_mask = (cols >= start) & (cols < end)
- block_mask = row_mask & col_mask # shape: (num_blocks, seq_len, seq_len)
-
- # Combine all blocks into one mask
- final_mask = torch.ones((seq_len, seq_len), dtype=self.config.dtype)
- final_mask[block_mask.any(dim=0)] = 0
-
- final_mask = torch.where(final_mask == 1.0, torch.finfo(q.dtype).min, final_mask)
-
- attention_mask[0] = final_mask
+ allowed = (row_mask & col_mask).any(dim=0)
+ blocked_mask = torch.full((seq_len, seq_len), torch.finfo(q.dtype).min, device=q.device, dtype=q.dtype)
+ final_mask = torch.where(allowed, torch.zeros_like(blocked_mask), blocked_mask)
+ attention_mask = final_mask.unsqueeze(0)
q = q.transpose(0, 1)
k = k.transpose(0, 1)
@@ -362,10 +392,11 @@ def eager_attention_forward(
value_states = repeat_kv(value, module.num_key_value_groups)
attn_weights = torch.matmul(query, key_states.transpose(2, 3)) / math.sqrt(module.head_dim)
+ mask_value = torch.full_like(attn_weights, MIN_MASKED_ATTENTION_VALUE, dtype=attn_weights.dtype)
+
if attention_mask is not None:
- attn_weights = torch.where(
- attention_mask, torch.tensor(MIN_MASKED_ATTENTION_VALUE, dtype=module.config.torch_dtype), attn_weights
- )
+ # Apply the attention mask
+ attn_weights = torch.where(attention_mask, mask_value, attn_weights)
attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query.dtype)
attn_output = torch.matmul(attn_weights, value_states)
@@ -519,8 +550,8 @@ def forward(
if output_attentions:
outputs += (self_attn_weights,)
- if use_cache:
- outputs += (present_key_value,)
+ # if use_cache:
+ # outputs += (present_key_value,)
return outputs
@@ -574,7 +605,7 @@ def forward(
)
hidden_states = inputs_embeds
- position_embeddings = self.rotary_emb(hidden_states, position_ids[1:])
+ position_embeddings = None
cos, sin = qeff_prepare_mrope_cos_sin(
self.cos_cached, self.sin_cached, position_ids[1:], self.config.rope_scaling["mrope_section"]
)
@@ -666,7 +697,7 @@ def get_submodules_for_export(self) -> Type[nn.Module]:
def forward(self, pixel_values, image_grid_thw):
image_embeds, deepstack_feature_lists = self.model.visual(pixel_values, grid_thw=image_grid_thw)
bs = image_grid_thw.shape[0]
- split_size = torch.floor_divide(torch.tensor(image_embeds.size(0)), bs)
+ split_size = image_embeds.shape[0] // bs
image_embeds = image_embeds.reshape(bs, split_size, image_embeds.size(1))
deepstack_features = torch.stack(
[feature.reshape(bs, split_size, feature.size(1)) for feature in deepstack_feature_lists],
@@ -704,6 +735,7 @@ def forward(
inputs_embeds = self.model.get_input_embeddings()(input_ids)
B, N, C = inputs_embeds.shape
selected = input_ids == self.model.config.image_token_id
+ # indices1 = qeff_cumsum_dim1(selected.to(torch.int64)) - 1
indices1 = selected.to(torch.int64).cumsum(1) - 1
indices1 = torch.where(indices1 != -1, indices1 + image_idx, indices1)
indices0 = torch.arange(selected.unsqueeze(0).shape[0]).view(-1, 1)
@@ -739,8 +771,15 @@ def forward(
logits = self.model.lm_head(hidden_states)
image_idx = (indices1.max() + 1).unsqueeze(0).unsqueeze(0)
if _should_export_embedding_output(self):
- return logits, vision_embeds, deepstack_features, image_idx, hidden_states, outputs.past_key_values
- return logits, vision_embeds, deepstack_features, image_idx, outputs.past_key_values
+ return (
+ logits,
+ vision_embeds.clone(),
+ deepstack_features.clone(),
+ image_idx,
+ hidden_states,
+ outputs.past_key_values,
+ )
+ return logits, vision_embeds.clone(), deepstack_features.clone(), image_idx, outputs.past_key_values
class QEffQwen3VLModel(Qwen3VLModel):
@@ -818,6 +857,7 @@ def forward(
inputs_embeds = self.model.get_input_embeddings()(input_ids)
B, N, C = inputs_embeds.shape
selected = input_ids == self.model.config.image_token_id
+ # indices1 = qeff_cumsum_dim1(selected.to(torch.int64)) - 1
indices1 = selected.to(torch.int64).cumsum(1) - 1
indices1 = torch.where(indices1 != -1, indices1 + image_idx, indices1)
indices0 = torch.arange(selected.unsqueeze(0).shape[0]).view(-1, 1)
@@ -848,64 +888,55 @@ def get_dummy_inputs(
continuous_batching: bool = False,
**kwargs,
):
+ bs: int = constants.ONNX_EXPORT_EXAMPLE_BATCH_SIZE + 1
+ fbs: int = constants.ONNX_EXPORT_EXAMPLE_FBS
prefill_seq_len = kwargs.get("prefill_seq_len", constants.ONNX_EXPORT_EXAMPLE_SEQ_LEN)
if prefill_seq_len is None:
prefill_seq_len = constants.ONNX_EXPORT_EXAMPLE_SEQ_LEN
prefill_seq_len = int(prefill_seq_len)
inputs_shapes = {}
- inputs_shapes["input_ids"] = (constants.ONNX_EXPORT_EXAMPLE_BATCH_SIZE, prefill_seq_len)
+ inputs_shapes["input_ids"] = (bs, prefill_seq_len)
# vision_size = 1024
vision_size = 187
inputs_shapes["vision_embeds"] = (
- constants.ONNX_EXPORT_EXAMPLE_BATCH_SIZE,
+ bs,
vision_size,
self.model.config.vision_config.out_hidden_size,
)
- inputs_shapes["image_grid_thw"] = (1, 1, 22, 34)
+ inputs_shapes["image_grid_thw"] = (bs, 1, 22, 34)
inputs_shapes["position_ids"] = (
3,
- constants.ONNX_EXPORT_EXAMPLE_BATCH_SIZE,
+ bs,
prefill_seq_len,
)
- inputs_shapes["pixel_values"] = (748, 1536)
+ inputs_shapes["pixel_values"] = (748 * bs, 1536)
inputs_shapes["image_idx"] = (1, 1)
- inputs_shapes["image_sizes"] = (constants.ONNX_EXPORT_EXAMPLE_BATCH_SIZE, 2)
+ inputs_shapes["image_sizes"] = (bs, 2)
inputs_shapes["deepstack_features"] = (
len(self.config.vision_config.deepstack_visual_indexes),
- constants.ONNX_EXPORT_EXAMPLE_BATCH_SIZE,
+ bs,
vision_size,
self.model.config.vision_config.out_hidden_size,
)
vision_inputs = {}
lang_inputs = {}
- vision_inputs["pixel_values"] = torch.zeros(
- (inputs_shapes["pixel_values"]), dtype=self.model.config.torch_dtype
- )
+ vision_inputs["pixel_values"] = torch.zeros((inputs_shapes["pixel_values"]), dtype=self.model.config.dtype)
vision_inputs["image_grid_thw"] = torch.zeros((inputs_shapes["image_grid_thw"]), dtype=torch.int64)
lang_inputs["input_ids"] = torch.zeros((inputs_shapes["input_ids"]), dtype=torch.int64)
- lang_inputs["vision_embeds"] = torch.zeros(
- (inputs_shapes["vision_embeds"]), dtype=self.model.config.torch_dtype
- )
+ lang_inputs["vision_embeds"] = torch.zeros((inputs_shapes["vision_embeds"]), dtype=self.model.config.dtype)
lang_inputs["position_ids"] = (
- (
- torch.arange(prefill_seq_len, dtype=torch.int64)
- .view(1, prefill_seq_len)
- .repeat(constants.ONNX_EXPORT_EXAMPLE_BATCH_SIZE, 1)
- )
+ (torch.arange(prefill_seq_len, dtype=torch.int64).view(1, prefill_seq_len).repeat(bs, 1))
.unsqueeze(0)
.repeat(4, 1, 1)
)
lang_inputs["image_idx"] = torch.zeros((inputs_shapes["image_idx"]), dtype=torch.int64)
lang_inputs["deepstack_features"] = torch.zeros(
- (inputs_shapes["deepstack_features"]), dtype=self.model.config.torch_dtype
+ (inputs_shapes["deepstack_features"]), dtype=self.model.config.dtype
)
# Add data for KV
- bs: int = constants.ONNX_EXPORT_EXAMPLE_BATCH_SIZE
- fbs: int = constants.ONNX_EXPORT_EXAMPLE_FBS
-
kv_cache_shape = get_padding_shape_from_config(
config=self.model.config.text_config,
batch_size=fbs if continuous_batching else bs,
@@ -915,9 +946,7 @@ def get_dummy_inputs(
lang_inputs["past_key_values"] = [[] for _ in range(self.model.config.text_config.num_hidden_layers)]
for i in range(self.model.config.text_config.num_hidden_layers):
for kv in ["key", "value"]:
- lang_inputs["past_key_values"][i].append(
- torch.zeros(kv_cache_shape, dtype=self.model.config.torch_dtype)
- )
+ lang_inputs["past_key_values"][i].append(torch.zeros(kv_cache_shape, dtype=self.model.config.dtype))
if continuous_batching:
lang_inputs["batch_index"] = torch.arange(bs).view(bs, 1)
@@ -1020,7 +1049,6 @@ def get_specializations(
"grid_h": grid_h,
"grid_w": grid_w,
"time": time,
- "num_feature_layers": len(self.config.vision_config.deepstack_visual_indexes),
}
)
@@ -1035,7 +1063,6 @@ def get_specializations(
"vision_size": max_vision_size,
"comp_ctx_lengths": comp_ctx_lengths_prefill[i],
"vision_batch_size": batch_size,
- "num_feature_layers": len(self.config.vision_config.deepstack_visual_indexes),
}
if continuous_batching:
@@ -1055,7 +1082,6 @@ def get_specializations(
"vision_size": max_vision_size,
"comp_ctx_lengths": comp_ctx_lengths_decode[i],
"vision_batch_size": batch_size,
- "num_feature_layers": len(self.config.vision_config.deepstack_visual_indexes),
}
if continuous_batching:
@@ -1071,7 +1097,6 @@ def get_specializations(
"ctx_len": ctx_len,
"vision_size": max_vision_size,
"vision_batch_size": batch_size,
- "num_feature_layers": len(self.config.vision_config.deepstack_visual_indexes),
}
if continuous_batching:
@@ -1087,7 +1112,6 @@ def get_specializations(
"ctx_len": ctx_len,
"vision_size": max_vision_size,
"vision_batch_size": batch_size,
- "num_feature_layers": len(self.config.vision_config.deepstack_visual_indexes),
}
if continuous_batching:
@@ -1114,16 +1138,16 @@ def get_onnx_dynamic_axes(
# Define dynamic axes
num_layers = self.config.text_config.num_hidden_layers
vision_dynamic_axes = {
- "pixel_values": {0: "grid_height", 1: "grid_width"},
- "image_grid_thw": {0: "batch_size", 1: "time", 2: "grid_h", 3: "grid_w"},
- "deepstack_features": {0: "num_feature_layers", 1: "batch_size", 2: "vision_size"},
+ "pixel_values": {0: "grid_height"},
+ "image_grid_thw": {0: "batch_size", 2: "grid_h", 3: "grid_w"},
+ "deepstack_features": {1: "batch_size", 2: "vision_size"},
}
lang_dynamic_axes = {
"input_ids": {0: "batch_size", 1: "seq_len"},
"position_ids": {1: "batch_size", 2: "seq_len"},
"vision_embeds": {0: "vision_batch_size", 1: "vision_size"},
- "deepstack_features": {0: "num_feature_layers", 1: "vision_batch_size", 2: "vision_size"},
+ "deepstack_features": {1: "vision_batch_size", 2: "vision_size"},
}
for i in range(num_layers):
@@ -1215,7 +1239,7 @@ def get_inputs_info(self):
IOInfo(name="attention_mask", datatype=torch.int64, shape=("batch_size", "seq_len")),
IOInfo(
name="pixel_values",
- datatype=self.config.torch_dtype,
+ datatype=self.config.dtype,
shape=("batch_size", 3, "image_size", "image_size"),
),
]
diff --git a/QEfficient/transformers/models/qwen3_vl_moe/modeling_qwen3_vl_moe.py b/QEfficient/transformers/models/qwen3_vl_moe/modeling_qwen3_vl_moe.py
index c36554ca38..563bb9a49c 100644
--- a/QEfficient/transformers/models/qwen3_vl_moe/modeling_qwen3_vl_moe.py
+++ b/QEfficient/transformers/models/qwen3_vl_moe/modeling_qwen3_vl_moe.py
@@ -77,6 +77,37 @@ def _batch_index_gather(tensor: torch.Tensor, batch_index: torch.Tensor) -> torc
return tensor.index_select(0, batch_index.reshape(-1).long())
+# def qeff_apply_interleaved_mrope(freqs, mrope_section):
+# """Apply interleaved MRoPE to 3D rotary embeddings.
+# Reorganizes frequency layout from chunked [TTT...HHH...WWW] to
+# interleaved [THWTHWTHW...TT], preserving frequency continuity.
+# args:
+# x: (3, bs, seq_len, head_dim // 2)
+# mrope_section: (3,)
+# returns:
+# x_t: (bs, seq_len, head_dim // 2)
+# """
+# freq_idx = torch.arange(freqs.shape[-1], device=freqs.device)
+# half_shape = freqs.shape[-1] // 2
+
+# h_mask = (freq_idx >= 1) & (freq_idx < mrope_section[1] * 3) & ((freq_idx - 1) % 3 == 0)
+# h_mask = h_mask | (
+# (freq_idx >= half_shape + 1)
+# & (freq_idx < half_shape + mrope_section[1] * 3)
+# & ((freq_idx - half_shape - 1) % 3 == 0)
+# )
+# w_mask = (freq_idx >= 2) & (freq_idx < mrope_section[2] * 3) & ((freq_idx - 2) % 3 == 0)
+# w_mask = w_mask | (
+# (freq_idx >= half_shape + 2)
+# & (freq_idx < half_shape + mrope_section[2] * 3)
+# & ((freq_idx - half_shape - 2) % 3 == 0)
+# )
+
+# freqs_t = torch.where(h_mask, freqs[1], freqs[0])
+# freqs_t = torch.where(w_mask, freqs[2], freqs_t)
+# return freqs_t
+
+
def qeff_apply_interleaved_mrope(freqs, mrope_section):
"""Apply interleaved MRoPE to 3D rotary embeddings.
Reorganizes frequency layout from chunked [TTT...HHH...WWW] to
@@ -87,15 +118,22 @@ def qeff_apply_interleaved_mrope(freqs, mrope_section):
returns:
x_t: (bs, seq_len, head_dim // 2)
"""
- freqs_t = freqs[0].clone()
+ freqs_t = freqs[0] # just overwrite the first dimension T
+ half_shape = freqs.shape[-1] // 2
for dim, offset in enumerate((1, 2), start=1): # H, W
length = mrope_section[dim] * 3
idx = slice(offset, length, 3)
freqs_t[..., idx] = freqs[dim, ..., idx]
+ offset += half_shape
+ length += half_shape
+ idx = slice(offset, length, 3)
+ freqs_t[..., idx] = freqs[dim, ..., idx]
return freqs_t
def qeff_prepare_mrope_cos_sin(cos, sin, position_ids, mrope_section, dtype=None):
+ cos = cos.to(device=position_ids.device)
+ sin = sin.to(device=position_ids.device)
invalid_pos_mask = position_ids < 0
safe_position_ids = torch.where(invalid_pos_mask, torch.zeros_like(position_ids), position_ids)
flat_pos = safe_position_ids.reshape(-1)
@@ -109,6 +147,12 @@ def qeff_prepare_mrope_cos_sin(cos, sin, position_ids, mrope_section, dtype=None
return cos, sin
+def qeff_cumsum_dim1(tensor: torch.Tensor) -> torch.Tensor:
+ seq_len = tensor.shape[1]
+ cumsum_mask = torch.tril(torch.ones((seq_len, seq_len), dtype=tensor.dtype, device=tensor.device))
+ return tensor @ cumsum_mask
+
+
def rotate_half_constant(x):
"""Rotates half the hidden dims of the input."""
_, _, _, hs = x.size()
@@ -154,9 +198,7 @@ class QEffQwen3VLMoeTextRotaryEmbedding(Qwen3VLMoeTextRotaryEmbedding):
def __init__(self, config: Qwen3VLMoeTextConfig, device=None):
super().__init__(config=config)
# Build here to make `torch.jit.trace` work.
- self._set_cos_sin_cache(
- seq_len=self.original_max_seq_len, device=self.inv_freq.device, dtype=config.torch_dtype
- )
+ self._set_cos_sin_cache(seq_len=self.original_max_seq_len, device=self.inv_freq.device, dtype=config.dtype)
def _set_cos_sin_cache(self, seq_len, device, dtype):
self.max_seq_len_cached = seq_len
@@ -172,14 +214,18 @@ def _set_cos_sin_cache(self, seq_len, device, dtype):
class QEffQwen3VLMoeVisionModel(Qwen3VLMoeVisionModel):
def rot_pos_emb(self, grid_thw: torch.Tensor) -> torch.Tensor:
merge_size = self.spatial_merge_size
- max_hw = max(grid_thw.shape)
- freq_table = self.rotary_pos_emb(max_hw) # (max_hw, dim // 2)
- device = freq_table.device
- bs, num_frames, height, width = grid_thw.shape
- grid_thw = (torch.tensor(grid_thw.shape, dtype=torch.int64)).unsqueeze(0)
+ # max_hw = max(grid_thw.shape)
+ # freq_table = self.rotary_pos_emb(max_hw) # (max_hw, dim // 2)
+ # device = freq_table.device
+ # bs, num_frames, height, width = grid_thw.shape
+ # grid_thw = (torch.tensor(grid_thw.shape, dtype=torch.int64)).unsqueeze(0)
- total_tokens = int(torch.prod(grid_thw, dim=1).sum().item())
- pos_ids = torch.empty((total_tokens, 2), dtype=torch.long, device=device)
+ # total_tokens = int(torch.prod(grid_thw, dim=1).sum().item())
+ # pos_ids = torch.empty((total_tokens, 2), dtype=torch.long, device=device)
+
+ bs, num_frames, height, width = grid_thw.shape
+ inv_freq = self.rotary_pos_emb.inv_freq
+ device = inv_freq.device
merged_h, merged_w = height // merge_size, width // merge_size
@@ -197,25 +243,28 @@ def rot_pos_emb(self, grid_thw: torch.Tensor) -> torch.Tensor:
coords = torch.stack((row_idx, col_idx), dim=-1)
- if num_frames > 1:
- coords = coords.repeat(num_frames, 1)
+ coords = coords.repeat(num_frames, 1)
- pos_ids = coords
- embeddings = freq_table[pos_ids] # lookup rotary embeddings
+ coords = coords.repeat(bs, 1)
+ embeddings = coords.to(dtype=inv_freq.dtype).unsqueeze(-1) * inv_freq.view(1, 1, -1)
embeddings = embeddings.flatten(1)
return embeddings
def fast_pos_embed_interpolate(self, grid_thw):
bs, t, h, w = grid_thw.shape
- h_idxs = torch.linspace(0, self.num_grid_per_side - 1, h)
- w_idxs = torch.linspace(0, self.num_grid_per_side - 1, w)
+ device = self.pos_embed.weight.device
+ h_den = torch.clamp(torch.scalar_tensor(h - 1, device=device, dtype=torch.float32), min=1.0)
+ w_den = torch.clamp(torch.scalar_tensor(w - 1, device=device, dtype=torch.float32), min=1.0)
+ h_idxs = torch.arange(h, device=device, dtype=torch.float32) * ((self.num_grid_per_side - 1) / h_den)
+ w_idxs = torch.arange(w, device=device, dtype=torch.float32) * ((self.num_grid_per_side - 1) / w_den)
h_idxs_floor = h_idxs.int()
w_idxs_floor = w_idxs.int()
- max_t = torch.tensor(self.num_grid_per_side - 1, device=h_idxs.device)
- h_idxs_ceil = torch.minimum(h_idxs_floor + 1, max_t) # working
- w_idxs_ceil = torch.minimum(w_idxs_floor + 1, max_t)
+ max_idx_h = torch.full_like(h_idxs_floor, self.num_grid_per_side - 1)
+ max_idx_w = torch.full_like(w_idxs_floor, self.num_grid_per_side - 1)
+ h_idxs_ceil = torch.minimum(h_idxs_floor + 1, max_idx_h)
+ w_idxs_ceil = torch.minimum(w_idxs_floor + 1, max_idx_w)
dh = h_idxs - h_idxs_floor
dw = w_idxs - w_idxs_floor
@@ -245,11 +294,8 @@ def fast_pos_embed_interpolate(self, grid_thw):
pos_embeds = self.pos_embed(idx_tensor) * weight_tensor[:, :, None]
patch_pos_embeds = pos_embeds[0] + pos_embeds[1] + pos_embeds[2] + pos_embeds[3]
- patch_pos_embeds = patch_pos_embeds.split([h * w])
-
- patch_pos_embeds_permute = []
merge_size = self.config.spatial_merge_size
- pos_embed = patch_pos_embeds[0]
+ pos_embed = patch_pos_embeds
pos_embed = pos_embed.repeat(t, 1)
pos_embed = (
@@ -257,14 +303,14 @@ def fast_pos_embed_interpolate(self, grid_thw):
.permute(0, 1, 3, 2, 4, 5)
.flatten(0, 4)
)
- patch_pos_embeds_permute.append(pos_embed)
- patch_pos_embeds = torch.cat(patch_pos_embeds_permute)
+ patch_pos_embeds = pos_embed
x_expanded = patch_pos_embeds.unsqueeze(0)
x_expanded = x_expanded.expand(bs, -1, -1)
patch_pos_embeds = x_expanded.reshape(-1, patch_pos_embeds.size(1))
return patch_pos_embeds
def forward(self, hidden_states: torch.Tensor, grid_thw: torch.Tensor) -> torch.Tensor:
+ grid_thw = grid_thw.to(device=hidden_states.device)
hidden_states = self.patch_embed(hidden_states)
pos_embeds = self.fast_pos_embed_interpolate(grid_thw)
hidden_states = hidden_states + pos_embeds
@@ -278,15 +324,15 @@ def forward(self, hidden_states: torch.Tensor, grid_thw: torch.Tensor) -> torch.
position_embeddings = (emb.cos(), emb.sin())
bs, t, h, w = grid_thw.shape
- t = torch.arange(t, t + 1).squeeze().expand(bs)
- h = torch.arange(h, h + 1).squeeze().expand(bs)
- w = torch.arange(w, w + 1).squeeze().expand(bs)
+ t = torch.arange(t, t + 1, device=grid_thw.device).squeeze().expand(bs)
+ h = torch.arange(h, h + 1, device=grid_thw.device).squeeze().expand(bs)
+ w = torch.arange(w, w + 1, device=grid_thw.device).squeeze().expand(bs)
cu_seqlens = (h * w).cumsum(
dim=0,
dtype=torch.int32,
)
- cu_seqlens = torch.cat([torch.tensor([0], dtype=cu_seqlens.dtype), cu_seqlens])
+ cu_seqlens = torch.cat([torch.zeros(1, dtype=cu_seqlens.dtype, device=cu_seqlens.device), cu_seqlens])
deepstack_feature_lists = []
for layer_num, blk in enumerate(self.blocks):
@@ -333,6 +379,8 @@ def forward(
sin = emb.sin()
else:
cos, sin = position_embeddings
+ cos = cos.to(device=q.device)
+ sin = sin.to(device=q.device)
q, k = apply_rotary_pos_emb_vision(q, k, cos, sin)
attention_mask = torch.full(
@@ -341,8 +389,8 @@ def forward(
# Create index grids
seq_len = attention_mask.shape[-1]
- rows = torch.arange(seq_len).view(1, -1)
- cols = torch.arange(seq_len).view(-1, 1)
+ rows = torch.arange(seq_len, device=q.device).view(1, -1)
+ cols = torch.arange(seq_len, device=q.device).view(-1, 1)
# Prepare start and end indices
start = cu_seqlens[:-1].view(-1, 1, 1)
@@ -353,13 +401,13 @@ def forward(
col_mask = (cols >= start) & (cols < end)
block_mask = row_mask & col_mask # shape: (num_blocks, seq_len, seq_len)
- # Combine all blocks into one mask
- final_mask = torch.ones((seq_len, seq_len), dtype=self.config.dtype)
- final_mask[block_mask.any(dim=0)] = 0
-
- final_mask = torch.where(final_mask == 1.0, torch.finfo(q.dtype).min, final_mask)
-
- attention_mask[0] = final_mask
+ allowed_mask = block_mask.any(dim=0)
+ final_mask = torch.full((seq_len, seq_len), torch.finfo(q.dtype).min, device=q.device, dtype=q.dtype)
+ attention_mask = torch.where(
+ allowed_mask.unsqueeze(0),
+ torch.zeros_like(attention_mask),
+ final_mask.unsqueeze(0),
+ )
q = q.transpose(0, 1)
k = k.transpose(0, 1)
@@ -390,9 +438,7 @@ def eager_attention_forward(
attn_weights = torch.matmul(query, key_states.transpose(2, 3)) / math.sqrt(module.head_dim)
if attention_mask is not None:
- attn_weights = torch.where(
- attention_mask, torch.tensor(MIN_MASKED_ATTENTION_VALUE, dtype=module.config.torch_dtype), attn_weights
- )
+ attn_weights = attn_weights.masked_fill(attention_mask, MIN_MASKED_ATTENTION_VALUE)
attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query.dtype)
attn_output = torch.matmul(attn_weights, value_states)
@@ -623,7 +669,6 @@ def forward(
)
hidden_states = inputs_embeds
- position_embeddings = self.rotary_emb(hidden_states, position_ids[1:])
cos, sin = qeff_prepare_mrope_cos_sin(
self.cos_cached,
self.sin_cached,
@@ -657,7 +702,6 @@ def forward(
output_attentions=output_attentions,
use_cache=use_cache,
cache_position=cache_position,
- position_embeddings=position_embeddings,
sin_cached=sin,
cos_cached=cos,
**kwargs,
@@ -775,7 +819,7 @@ def get_submodules_for_export(self) -> Type[nn.Module]:
def forward(self, pixel_values, image_grid_thw):
image_embeds, deepstack_feature_lists = self.model.visual(pixel_values, grid_thw=image_grid_thw)
bs = image_grid_thw.shape[0]
- split_size = torch.floor_divide(torch.tensor(image_embeds.size(0)), bs)
+ split_size = image_embeds.shape[0] // bs
image_embeds = image_embeds.reshape(bs, split_size, image_embeds.size(1))
deepstack_features = torch.stack(
[feature.reshape(bs, split_size, feature.size(1)) for feature in deepstack_feature_lists],
@@ -883,16 +927,19 @@ def forward(
# a single forward, identical to the pre-layerwise behavior/output contract.
B, N, C = inputs_embeds.shape
selected = input_ids == self.model.config.image_token_id
+ # indices1 = qeff_cumsum_dim1(selected.to(torch.int64)) - 1
indices1 = selected.to(torch.int64).cumsum(1) - 1
indices1 = torch.where(indices1 != -1, indices1 + image_idx, indices1)
- indices0 = torch.arange(selected.unsqueeze(0).shape[0]).view(-1, 1)
+ indices0 = torch.arange(selected.unsqueeze(0).shape[0], device=selected.device).view(-1, 1)
image_features_expanded = vision_embeds.reshape(-1, C).unsqueeze(0)[indices0, indices1]
num_features, bs, split_size, C = deepstack_features.shape
x = deepstack_features.reshape(num_features, bs * split_size, C)
deepstack_features_expanded = x[:, indices1, :]
image_input_embeds = torch.where(selected.unsqueeze(-1), image_features_expanded, inputs_embeds)
- inputs_embeds = torch.where(input_ids.shape[1] == torch.tensor(1), inputs_embeds, image_input_embeds)
+ inputs_embeds = torch.where(
+ input_ids.shape[1] == torch.tensor(1, device=input_ids.device), inputs_embeds, image_input_embeds
+ )
image_mask = selected.clone()
visual_pos_masks = None
@@ -912,26 +959,31 @@ def forward(
deepstack_visual_embeds=deepstack_visual_embeds,
)
logit_index = position_ids[0].to(torch.int32).argmax(1, keepdim=True)
- hidden_states = outputs.last_hidden_state[torch.arange(position_ids[0].shape[0]).view(-1, 1), logit_index]
+ hidden_states = outputs.last_hidden_state[
+ torch.arange(position_ids[0].shape[0], device=position_ids.device).view(-1, 1), logit_index
+ ]
if batch_fold_cb:
hidden_states = _batch_index_gather(hidden_states, batch_index)
logits = self.model.lm_head(hidden_states)
image_idx = (indices1.max() + 1).unsqueeze(0).unsqueeze(0)
- return logits, vision_embeds, deepstack_features, image_idx, outputs.past_key_values
+ return logits, vision_embeds.clone(), deepstack_features.clone(), image_idx, outputs.past_key_values
if QEffQwen3VLMoeTextModel._start == 0:
B, N, C = inputs_embeds.shape
selected = input_ids == self.model.config.image_token_id
+ # indices1 = qeff_cumsum_dim1(selected.to(torch.int64)) - 1
indices1 = selected.to(torch.int64).cumsum(1) - 1
indices1 = torch.where(indices1 != -1, indices1 + image_idx, indices1)
- indices0 = torch.arange(selected.unsqueeze(0).shape[0]).view(-1, 1)
+ indices0 = torch.arange(selected.unsqueeze(0).shape[0], device=selected.device).view(-1, 1)
image_features_expanded = vision_embeds.reshape(-1, C).unsqueeze(0)[indices0, indices1]
num_features, bs, split_size, C = deepstack_features.shape
x = deepstack_features.reshape(num_features, bs * split_size, C)
deepstack_features_expanded = x[:, indices1, :]
image_input_embeds = torch.where(selected.unsqueeze(-1), image_features_expanded, inputs_embeds)
- inputs_embeds = torch.where(input_ids.shape[1] == torch.tensor(1), inputs_embeds, image_input_embeds)
+ inputs_embeds = torch.where(
+ input_ids.shape[1] == torch.tensor(1, device=input_ids.device), inputs_embeds, image_input_embeds
+ )
image_mask = selected.clone()
@@ -960,7 +1012,7 @@ def forward(
hidden_states = outputs.last_hidden_state[:, -1:, :]
logits = hidden_states
image_idx = (indices1.max() + 1).unsqueeze(0).unsqueeze(0)
- return logits, vision_embeds, deepstack_features, image_idx, outputs.past_key_values
+ return logits, vision_embeds.clone(), deepstack_features.clone(), image_idx, outputs.past_key_values
elif QEffQwen3VLMoeTextModel._end == QEffQwen3VLMoeTextModel._total_layers:
outputs = self.language_model(
@@ -974,7 +1026,9 @@ def forward(
deepstack_visual_embeds=QEffQwen3VLDecoderWrapper._deepstack,
)
logit_index = position_ids[0].to(torch.int32).argmax(1, keepdim=True)
- hidden_states = outputs.last_hidden_state[torch.arange(position_ids[0].shape[0]).view(-1, 1), logit_index]
+ hidden_states = outputs.last_hidden_state[
+ torch.arange(position_ids[0].shape[0], device=position_ids.device).view(-1, 1), logit_index
+ ]
if batch_fold_cb:
hidden_states = _batch_index_gather(hidden_states, batch_index)
logits = self.model.lm_head(hidden_states)
@@ -1055,13 +1109,8 @@ def get_dummy_inputs(
continuous_batching: bool = False,
**kwargs,
):
- bs = kwargs.get("batch_size", constants.ONNX_EXPORT_EXAMPLE_BATCH_SIZE)
- if bs > 1:
- bs = 2
+ bs: int = constants.ONNX_EXPORT_EXAMPLE_BATCH_SIZE + 1
fbs: int = constants.ONNX_EXPORT_EXAMPLE_FBS
- batch_fold = kwargs.pop("batch_fold", False)
- if continuous_batching and batch_fold:
- bs = fbs
prefill_seq_len = kwargs.get("prefill_seq_len")
if prefill_seq_len is None:
@@ -1076,13 +1125,13 @@ def get_dummy_inputs(
vision_size,
self.model.config.vision_config.out_hidden_size,
)
- inputs_shapes["image_grid_thw"] = (1, 1, 22, 34)
+ inputs_shapes["image_grid_thw"] = (bs, 1, 22, 34)
inputs_shapes["position_ids"] = (
3,
bs,
prefill_seq_len,
)
- inputs_shapes["pixel_values"] = (748, 1536)
+ inputs_shapes["pixel_values"] = (748 * bs, 1536)
inputs_shapes["image_idx"] = (1, 1)
inputs_shapes["image_sizes"] = (bs, 2)
inputs_shapes["deepstack_features"] = (
@@ -1094,14 +1143,10 @@ def get_dummy_inputs(
vision_inputs = {}
lang_inputs = {}
- vision_inputs["pixel_values"] = torch.zeros(
- (inputs_shapes["pixel_values"]), dtype=self.model.config.torch_dtype
- )
+ vision_inputs["pixel_values"] = torch.zeros((inputs_shapes["pixel_values"]), dtype=self.model.config.dtype)
vision_inputs["image_grid_thw"] = torch.zeros((inputs_shapes["image_grid_thw"]), dtype=torch.int64)
lang_inputs["input_ids"] = torch.zeros((inputs_shapes["input_ids"]), dtype=torch.int64)
- lang_inputs["vision_embeds"] = torch.zeros(
- (inputs_shapes["vision_embeds"]), dtype=self.model.config.torch_dtype
- )
+ lang_inputs["vision_embeds"] = torch.zeros((inputs_shapes["vision_embeds"]), dtype=self.model.config.dtype)
lang_inputs["position_ids"] = (
(torch.arange(prefill_seq_len, dtype=torch.int64).view(1, prefill_seq_len).repeat(bs, 1))
.unsqueeze(0)
@@ -1109,7 +1154,7 @@ def get_dummy_inputs(
)
lang_inputs["image_idx"] = torch.zeros((inputs_shapes["image_idx"]), dtype=torch.int64)
lang_inputs["deepstack_features"] = torch.zeros(
- (inputs_shapes["deepstack_features"]), dtype=self.model.config.torch_dtype
+ (inputs_shapes["deepstack_features"]), dtype=self.model.config.dtype
)
# Add data for KV
@@ -1122,9 +1167,7 @@ def get_dummy_inputs(
lang_inputs["past_key_values"] = [[] for _ in range(self.model.config.text_config.num_hidden_layers)]
for i in range(self.model.config.text_config.num_hidden_layers):
for kv in ["key", "value"]:
- lang_inputs["past_key_values"][i].append(
- torch.zeros(kv_cache_shape, dtype=self.model.config.torch_dtype)
- )
+ lang_inputs["past_key_values"][i].append(torch.zeros(kv_cache_shape, dtype=self.model.config.dtype))
if continuous_batching:
lang_inputs["batch_index"] = torch.arange(bs).view(bs, 1)
@@ -1238,7 +1281,6 @@ def get_specializations(
"grid_h": grid_h,
"grid_w": grid_w,
"time": time,
- "num_feature_layers": len(self.config.vision_config.deepstack_visual_indexes),
}
)
@@ -1253,7 +1295,6 @@ def get_specializations(
"vision_size": vision_size,
"comp_ctx_lengths": comp_ctx_lengths_prefill[i],
"vision_batch_size": batch_size,
- "num_feature_layers": len(self.config.vision_config.deepstack_visual_indexes),
}
if continuous_batching:
@@ -1273,7 +1314,6 @@ def get_specializations(
"vision_size": vision_size,
"comp_ctx_lengths": comp_ctx_lengths_decode[i],
"vision_batch_size": batch_size,
- "num_feature_layers": len(self.config.vision_config.deepstack_visual_indexes),
}
if continuous_batching:
@@ -1289,7 +1329,6 @@ def get_specializations(
"ctx_len": ctx_len,
"vision_size": vision_size,
"vision_batch_size": batch_size,
- "num_feature_layers": len(self.config.vision_config.deepstack_visual_indexes),
}
if continuous_batching:
@@ -1305,7 +1344,6 @@ def get_specializations(
"ctx_len": ctx_len,
"vision_size": vision_size,
"vision_batch_size": batch_size,
- "num_feature_layers": len(self.config.vision_config.deepstack_visual_indexes),
}
if continuous_batching:
@@ -1337,16 +1375,16 @@ def get_onnx_dynamic_axes(
num_layers = self.config.text_config.num_hidden_layers
batch_axis = "full_batch_size" if continuous_batching and batch_fold else "batch_size"
vision_dynamic_axes = {
- "pixel_values": {0: "grid_height", 1: "grid_width"},
- "image_grid_thw": {0: "batch_size", 1: "time", 2: "grid_h", 3: "grid_w"},
- "deepstack_features": {0: "num_feature_layers", 1: "batch_size", 2: "vision_size"},
+ "pixel_values": {0: "grid_height"},
+ "image_grid_thw": {0: "batch_size", 2: "grid_h", 3: "grid_w"},
+ "deepstack_features": {1: "batch_size", 2: "vision_size"},
}
lang_dynamic_axes = {
"input_ids": {0: batch_axis, 1: "seq_len"},
"position_ids": {1: batch_axis, 2: "seq_len"},
"vision_embeds": {0: "vision_batch_size", 1: "vision_size"},
- "deepstack_features": {0: "num_feature_layers", 1: "vision_batch_size", 2: "vision_size"},
+ "deepstack_features": {1: "vision_batch_size", 2: "vision_size"},
}
for i in range(num_layers):
@@ -1398,7 +1436,9 @@ def get_output_names(self, kv_offload: bool = False):
def prepare_inputs_for_generation(self, inputs, prefill_seq_len=128, batch_size=1):
input_ids_length = inputs["input_ids"].shape[1]
- inputs["position_ids"] = torch.arange(input_ids_length).view(1, 1, input_ids_length).expand(-1, batch_size, -1)
+ inputs["position_ids"] = torch.arange(input_ids_length, device=inputs["input_ids"].device).view(
+ 1, 1, input_ids_length
+ ).expand(-1, batch_size, -1)
mm_token_type_ids = inputs.get("mm_token_type_ids")
if mm_token_type_ids is None:
@@ -1433,7 +1473,7 @@ def get_inputs_info(self):
IOInfo(name="attention_mask", datatype=torch.int64, shape=("batch_size", "seq_len")),
IOInfo(
name="pixel_values",
- datatype=self.config.torch_dtype,
+ datatype=self.config.dtype,
shape=("batch_size", 3, "image_size", "image_size"),
),
]
diff --git a/QEfficient/utils/export_utils.py b/QEfficient/utils/export_utils.py
index 78bf2675be..407742b1d3 100644
--- a/QEfficient/utils/export_utils.py
+++ b/QEfficient/utils/export_utils.py
@@ -12,14 +12,16 @@
from contextlib import contextmanager, nullcontext
from contextvars import ContextVar
from pathlib import Path
-from typing import Any, Dict
+from typing import Any
import torch
-import torch.nn as nn
+from torch import nn
from torch.export import Dim
from QEfficient.base.onnx_transforms import (
+ CanonicalizeWhileLoopInitialConditionTransform,
CustomOpTransform,
+ DeduplicateRepeatedSubgraphTransform,
PreserveNestedCacheRetainedStateTransform,
RenameFunctionOutputsTransform,
RenameRepeatedSubgraphTransform,
@@ -28,10 +30,10 @@
from QEfficient.transformers.cache_utils import InvalidIndexProvider
from QEfficient.utils.cache import QEFF_HOME
from QEfficient.utils.constants import (
- _KNOWN_DECODER_LAYER_ATTR_PATHS,
- _KNOWN_DECODER_LAYER_SUFFIXES,
DYNAMO_DIM_MAX_BATCH_SIZE,
DYNAMO_DIM_MIN_COMP_CTX_LENGTHS,
+ _KNOWN_DECODER_LAYER_ATTR_PATHS,
+ _KNOWN_DECODER_LAYER_SUFFIXES,
)
from QEfficient.utils.hash_utils import create_export_hash
from QEfficient.utils.logging_utils import logger
@@ -59,20 +61,20 @@ def reorder_inputs_by_signature(model, example_inputs, dynamic_shapes=None):
"""Reorder example_inputs (and optional dynamic_shapes) to match model.forward signature.
torch.export requires inputs and dynamic_shapes to follow the forward parameter order
- so that each shape constraint binds to the correct input tensor.
+ so that each shape constraint binds to the correct input tensor. Non-input
+ dynamic_shapes entries are dropped because torch.export only accepts entries
+ matching real forward inputs.
"""
sig_keys = list(inspect.signature(model.forward).parameters.keys())
sig_key_set = set(sig_keys)
- ordered_inputs, ordered_shapes = {}, {}
+ ordered_inputs = {}
for k in sig_keys:
if k in example_inputs:
ordered_inputs[k] = example_inputs[k]
- if dynamic_shapes is not None and k in dynamic_shapes:
- ordered_shapes[k] = dynamic_shapes[k]
reordered_inputs = {**ordered_inputs, **{k: v for k, v in example_inputs.items() if k not in sig_key_set}}
if dynamic_shapes is not None:
- reordered_shapes = {**ordered_shapes, **{k: v for k, v in dynamic_shapes.items() if k not in sig_key_set}}
- return reordered_inputs, reordered_shapes
+ dynamic_shapes = {k: dynamic_shapes.get(k) for k in reordered_inputs}
+ return reordered_inputs, dynamic_shapes
return reordered_inputs, None
@@ -98,9 +100,9 @@ def build_dynamo_export_kwargs(export_kwargs):
def convert_dynamic_axes_to_dynamic_shapes(
- dynamic_axes: Dict[str, Dict[int, str]],
+ dynamic_axes: dict[str, dict[int, str]],
model_config=None,
-) -> Dict[str, Any]:
+) -> dict[str, Any]:
"""
Convert ONNX dynamic_axes format to torch.export dynamic_shapes format.
@@ -125,40 +127,18 @@ def convert_dynamic_axes_to_dynamic_shapes(
torch.export dynamic_shapes dict with Dim objects, suitable for
torch.onnx.export(dynamic_shapes=...).
"""
- max_seq_len = getattr(model_config, "max_position_embeddings", 1024)
- model_type = getattr(model_config, "model_type", None)
- batch_min = 1 if model_type == "gpt_oss" else 2
-
- dim_registry: Dict[str, Any] = {}
+ dim_registry: dict[str, Any] = {}
def resolve_dim(dim_name: str):
if dim_name not in dim_registry:
- if dim_name == "batch_size":
- dim_registry[dim_name] = Dim("batch_size", min=batch_min, max=DYNAMO_DIM_MAX_BATCH_SIZE)
- elif dim_name == "full_batch_size":
- # CB pool capacity; different min prevents torch.export collapsing it with batch_size.
- dim_registry[dim_name] = Dim("full_batch_size", min=batch_min + 1, max=DYNAMO_DIM_MAX_BATCH_SIZE)
- elif "seq_len" in dim_name:
- dim_registry[dim_name] = Dim("seq_len", min=2, max=max_seq_len)
- elif "comp_ctx_lengths" in dim_name:
- dim_registry[dim_name] = Dim("comp_ctx_lengths", min=DYNAMO_DIM_MIN_COMP_CTX_LENGTHS, max=max_seq_len)
- elif "ctx_len" in dim_name:
- dim_registry[dim_name] = Dim("ctx_len", min=2, max=max_seq_len)
- elif "sliding_window" in dim_name:
- dim_registry[dim_name] = Dim(
- "sliding_window",
- min=2,
- max=getattr(model_config, "sliding_window", max_seq_len),
- )
- else:
- dim_registry[dim_name] = Dim.DYNAMIC
+ dim_registry[dim_name] = Dim(dim_name)
return dim_registry[dim_name]
- dynamic_shapes: Dict[str, Any] = {}
- past_keys: Dict[int, Any] = {}
- past_values: Dict[int, Any] = {}
- compressed_kv_layers: Dict[int, Any] = {}
- k_pe_layers: Dict[int, Any] = {}
+ dynamic_shapes: dict[str, Any] = {}
+ past_keys: dict[int, Any] = {}
+ past_values: dict[int, Any] = {}
+ compressed_kv_layers: dict[int, Any] = {}
+ k_pe_layers: dict[int, Any] = {}
for input_name, axes_map in dynamic_axes.items():
resolved = {axis_idx: resolve_dim(dim_name) for axis_idx, dim_name in axes_map.items()}
@@ -187,6 +167,50 @@ def resolve_dim(dim_name: str):
return dynamic_shapes
+ dynamic_shapes: dict[str, Any] = {}
+ past_keys: dict[int, Any] = {}
+ past_values: dict[int, Any] = {}
+ hybrid_states: dict[int, list[Any]] = {}
+ compressed_kv_layers: dict[int, Any] = {}
+ k_pe_layers: dict[int, Any] = {}
+
+ for input_name, axes_map in dynamic_axes.items():
+ resolved = {}
+ for axis_idx, dim_name in axes_map.items():
+ dim = resolve_dim(dim_name)
+ if isinstance(dim, int):
+ continue
+ resolved[axis_idx] = dim
+ if input_name.startswith("past_key."):
+ past_keys[int(input_name.split(".")[1])] = resolved
+ elif input_name.startswith("past_value."):
+ past_values[int(input_name.split(".")[1])] = resolved
+ elif input_name.startswith("conv_state."):
+ hybrid_states.setdefault(int(input_name.split(".")[1]), [{}, {}])[0] = resolved
+ elif input_name.startswith("recurrent_state."):
+ hybrid_states.setdefault(int(input_name.split(".")[1]), [{}, {}])[1] = resolved
+ elif input_name.startswith("compressed_kv."):
+ compressed_kv_layers[int(input_name.split(".")[1])] = resolved
+ elif input_name.startswith("k_pe."):
+ k_pe_layers[int(input_name.split(".")[1])] = resolved
+ else:
+ dynamic_shapes[input_name] = resolved
+
+ if past_keys or past_values or hybrid_states:
+ max_layer = max(list(past_keys.keys()) + list(past_values.keys()) + list(hybrid_states.keys()))
+ dynamic_shapes["past_key_values"] = [
+ hybrid_states.get(i, [past_keys.get(i, {}), past_values.get(i, {})])
+ for i in range(max_layer + 1)
+ ]
+
+ if compressed_kv_layers or k_pe_layers:
+ max_layer = max(list(compressed_kv_layers.keys()) + list(k_pe_layers.keys()))
+ dynamic_shapes["compressed_kvs"] = [
+ (compressed_kv_layers.get(i, {}), k_pe_layers.get(i, {})) for i in range(max_layer + 1)
+ ]
+
+ return dynamic_shapes
+
def _resolve_attr_path(root, attr_path):
current = root
@@ -342,9 +366,8 @@ def wrapper(self, *args, **kwargs):
else nullcontext()
)
try:
- with export_context:
- with dynamo_patch:
- onnx_path = func(self, *args, **kwargs)
+ with export_context, dynamo_patch:
+ onnx_path = func(self, *args, **kwargs)
except Exception as export_exc:
if use_onnx_subfunctions and dynamo:
raise RuntimeError(
@@ -519,8 +542,16 @@ def _setup_onnx_subfunctions(qeff_model, args, kwargs, dynamo=False):
# Dynamo: PreserveNestedCacheRetainedStateTransform + RenameRepeatedSubgraphTransform.
if PreserveNestedCacheRetainedStateTransform not in qeff_model._onnx_transforms:
qeff_model._onnx_transforms.append(PreserveNestedCacheRetainedStateTransform)
+ if CanonicalizeWhileLoopInitialConditionTransform not in qeff_model._onnx_transforms:
+ qeff_model._onnx_transforms.append(CanonicalizeWhileLoopInitialConditionTransform)
+ if DeduplicateRepeatedSubgraphTransform not in qeff_model._onnx_transforms:
+ qeff_model._onnx_transforms.append(DeduplicateRepeatedSubgraphTransform)
if RenameRepeatedSubgraphTransform not in qeff_model._onnx_transforms:
qeff_model._onnx_transforms.append(RenameRepeatedSubgraphTransform)
+ # if QualifyOnnxNodeNamesTransform not in qeff_model._onnx_transforms:
+ # qeff_model._onnx_transforms.append(QualifyOnnxNodeNamesTransform)
+ # if RewriteSequenceSplitGetItemTransform not in qeff_model._onnx_transforms:
+ # qeff_model._onnx_transforms.append(RewriteSequenceSplitGetItemTransform)
else:
# TorchScript: RenameFunctionOutputsTransform + CustomOpTransform.
if RenameFunctionOutputsTransform not in qeff_model._onnx_transforms:
@@ -539,6 +570,9 @@ def _setup_onnx_subfunctions(qeff_model, args, kwargs, dynamo=False):
qeff_model._subfunction_target_classnames = resolved_classnames
onnx_transform_kwargs = dict(kwargs.get("onnx_transform_kwargs") or {})
onnx_transform_kwargs["target_classnames"] = resolved_classnames
+ onnx_transform_kwargs["target_class_modules"] = {
+ cls.__name__: cls.__module__ for cls in decoder_layer_classes
+ }
kwargs["onnx_transform_kwargs"] = onnx_transform_kwargs
else:
# TorchScript path: pass class objects for export_modules_as_functions
@@ -590,7 +624,7 @@ def _cleanup_onnx_subfunctions(qeff_model, state=None):
qeff_model.hash_params["onnx_subfunction_version"] = state["hash_subfunction_version"]
-def _save_export_metadata(export_dir: Path, filtered_hash_params: Dict):
+def _save_export_metadata(export_dir: Path, filtered_hash_params: dict):
"""
Save export metadata to JSON file for reproducibility.
diff --git a/examples/image_text_to_text/models/qwen3_5/qwen3_5.py b/examples/image_text_to_text/models/qwen3_5/qwen3_5.py
index 872b50efe3..e5541c67c4 100644
--- a/examples/image_text_to_text/models/qwen3_5/qwen3_5.py
+++ b/examples/image_text_to_text/models/qwen3_5/qwen3_5.py
@@ -5,26 +5,50 @@
#
# -----------------------------------------------------------------------------
+import os
+
+import numpy as np
import requests
+import torch
import transformers
from PIL import Image
from qwen_vl_utils import process_vision_info
-from transformers import AutoConfig, AutoProcessor, TextStreamer
+from transformers import AutoConfig, AutoProcessor
from QEfficient import QEFFAutoModelForImageTextToText
+from QEfficient.generation.cloud_infer import QAICInferenceSession
model_id = "Qwen/Qwen3.5-0.8B"
+DECODE_NUM_DEVICES = int(os.environ.get("QEFF_DECODE_NUM_DEVICES", "1"))
+WEIGHT_FREE = os.environ.get("QEFF_WEIGHT_FREE", "0") == "1"
config = AutoConfig.from_pretrained(model_id)
# For faster execution user can run with lesser layers, For Testing Purpose Only
# config.vision_config.depth = 4
# config.text_config.num_hidden_layers = 4
config.torch_dtype = "float32"
+layer_types = list(getattr(config.text_config, "layer_types", []))
+if len(layer_types) < config.text_config.num_hidden_layers:
+ layer_types.extend(["full_attention"] * (config.text_config.num_hidden_layers - len(layer_types)))
+config.text_config.layer_types = layer_types[: config.text_config.num_hidden_layers]
+
+
+def _update_retained_states(target_inputs, source_outputs):
+ for layer_idx, layer_type in enumerate(config.text_config.layer_types):
+ # if layer_type == "full_attention":
+ # state_names = (f"past_key.{layer_idx}", f"past_value.{layer_idx}")
+ # else:
+ # state_names = (f"conv_state.{layer_idx}", f"recurrent_state.{layer_idx}")
+
+ state_names = (f"past_key.{layer_idx}", f"past_value.{layer_idx}")
+ for state_name in state_names:
+ target_inputs[state_name] = source_outputs[f"{state_name}_RetainedState"]
qeff_model = QEFFAutoModelForImageTextToText.from_pretrained(
model_id,
attn_implementation="eager",
kv_offload=True,
+ weight_free=WEIGHT_FREE,
config=config,
# # For CCL activation
# qaic_config={
@@ -56,24 +80,47 @@
if skip_vision:
## Only Text ##
- qeff_model.compile(
+ prefill_qpc_path = qeff_model.compile(
batch_size=BS,
prefill_seq_len=PREFILL_SEQ_LEN,
ctx_len=CTX_LEN,
num_cores=16,
- num_devices=4,
+ num_devices=1,
mxfp6_matmul=True,
mxint8_kv_cache=True,
+ retain_full_kv=True,
+ split_model_io=True,
aic_enable_depth_first=False,
+ prefill_only=True,
+ enable_chunking=True,
skip_vision=True,
mos=1,
- split_model_io=True,
use_onnx_subfunctions=True,
+ dynamo=True,
# comp_ctx_lengths_prefill=comp_ctx_lengths_prefill,
# comp_ctx_lengths_decode=comp_ctx_lengths_decode,
# qaic_config=qaic_config, # Enable KV blocking - comment out to disable
)
+ decode_qpc_path = qeff_model.compile(
+ batch_size=BS,
+ prefill_seq_len=1,
+ ctx_len=CTX_LEN,
+ num_cores=16,
+ num_devices=DECODE_NUM_DEVICES,
+ mxfp6_matmul=True,
+ mxint8_kv_cache=True,
+ retain_full_kv=True,
+ split_model_io=True,
+ aic_enable_depth_first=True,
+ prefill_only=False,
+ skip_vision=True,
+ mos=1,
+ use_onnx_subfunctions=True,
+ dynamo=True,
+ # qaic_config=qaic_config, # Enable KV blocking - comment out to disable
+ )
+
if enable_blocking:
print("\n" + "=" * 80)
print("Verifying KV Blocking Applied During Compilation")
@@ -96,39 +143,29 @@
print(" Status: INACTIVE - Model compiled without blocking")
print("=" * 80 + "\n")
- text_prompt_2 = "Describe yourself as a large language model, including your purpose, capabilities, and limitations. Explain how you process and generate responses, interact with users, and handle uncertainty, while emphasizing accuracy, safety, and helpfulness in diverse conversations across various topics and domains."
-
- messages = [
- {
- "role": "user",
- "content": [
- {"type": "text", "text": text_prompt_2},
- ],
- },
- ]
-
- messages = [messages] * BS
-
- inputs = processor.apply_chat_template(
- messages,
- add_generation_prompt=True,
- tokenize=True,
- return_dict=True,
- return_tensors="pt",
- )
- inputs = qeff_model.model.prepare_inputs_for_generation(
- inputs=inputs, prefill_seq_len=PREFILL_SEQ_LEN, batch_size=BS
- )
- streamer = TextStreamer(tokenizer)
- output = qeff_model.generate(inputs=inputs, generation_len=512, streamer=streamer)
- print(output.generated_ids)
- print(tokenizer.batch_decode(output.generated_ids))
- print(output)
-
else:
## Vision + Text ##
- qeff_model.compile(
+ vision_qpc_path = qeff_model.compile(
+ batch_size=BS,
+ prefill_seq_len=PREFILL_SEQ_LEN,
+ ctx_len=CTX_LEN,
+ num_cores=16,
+ num_devices=1,
+ height=354,
+ width=536,
+ mxfp6_matmul=True,
+ mxint8_kv_cache=True,
+ aic_enable_depth_first=False,
+ mos=1,
+ split_model_io=True,
+ skip_vision=False,
+ skip_lang=True,
+ use_onnx_subfunctions=True,
+ dynamo=True,
+ )
+
+ prefill_qpc_path = qeff_model.compile(
batch_size=BS,
prefill_seq_len=PREFILL_SEQ_LEN,
ctx_len=CTX_LEN,
@@ -138,10 +175,36 @@
width=536,
mxfp6_matmul=True,
mxint8_kv_cache=True,
+ retain_full_kv=True,
+ split_model_io=True,
aic_enable_depth_first=False,
+ prefill_only=True,
+ enable_chunking=True,
+ skip_vision=True,
mos=1,
+ use_onnx_subfunctions=True,
+ dynamo=True,
+ # qaic_config=qaic_config, # Enable KV blocking - comment out to disable
+ )
+
+ decode_qpc_path = qeff_model.compile(
+ batch_size=BS,
+ prefill_seq_len=1,
+ ctx_len=CTX_LEN,
+ num_cores=16,
+ num_devices=DECODE_NUM_DEVICES,
+ height=354,
+ width=536,
+ mxfp6_matmul=True,
+ mxint8_kv_cache=True,
+ retain_full_kv=True,
split_model_io=True,
+ aic_enable_depth_first=True,
+ prefill_only=False,
+ skip_vision=True,
+ mos=1,
use_onnx_subfunctions=True,
+ dynamo=True,
# comp_ctx_lengths_prefill=comp_ctx_lengths_prefill,
# comp_ctx_lengths_decode=comp_ctx_lengths_decode,
# qaic_config=qaic_config, # Enable KV blocking - comment out to disable
@@ -196,11 +259,119 @@
padding=True,
return_tensors="pt",
)
- inputs = qeff_model.model.prepare_inputs_for_generation(
- inputs=inputs, prefill_seq_len=PREFILL_SEQ_LEN, batch_size=BS
+ inputs = qeff_model.model.prepare_inputs_for_generation(inputs=inputs, prefill_seq_len=PREFILL_SEQ_LEN, batch_size=BS)
+
+lang_prefill_session = QAICInferenceSession(prefill_qpc_path.get("lang_prefill_qpc_path"))
+lang_decode_session = QAICInferenceSession(decode_qpc_path.get("lang_decode_qpc_path"))
+vision_session = None
+if not skip_vision:
+ vision_session = QAICInferenceSession(vision_qpc_path.get("vision_qpc_path"))
+
+if skip_vision:
+ messages = [
+ {
+ "role": "user",
+ "content": [
+ {
+ "type": "text",
+ "text": "Describe yourself as a large language model, including your purpose, capabilities, and limitations. Explain how you process and generate responses, interact with users, and handle uncertainty, while emphasizing accuracy, safety, and helpfulness in diverse conversations across various topics and domains.",
+ },
+ ],
+ },
+ ]
+else:
+ messages = [
+ {
+ "role": "user",
+ "content": [
+ {"type": "image", "image": image},
+ {"type": "text", "text": "Describe all the colors seen in the image."},
+ ],
+ },
+ ]
+
+messages = [messages] * BS
+texts = [processor.apply_chat_template(msg, tokenize=False, add_generation_prompt=True) for msg in messages]
+image_inputs, video_inputs = process_vision_info(messages)
+inputs = processor(text=texts, images=image_inputs, videos=video_inputs, padding=True, return_tensors="pt")
+inputs = qeff_model.model.prepare_inputs_for_generation(inputs=inputs, prefill_seq_len=PREFILL_SEQ_LEN, batch_size=BS)
+
+pad_token_id = tokenizer.pad_token_id
+input_ids_length = inputs["input_ids"].shape[1]
+num_chunks = -(input_ids_length // -PREFILL_SEQ_LEN)
+padded_len = num_chunks * PREFILL_SEQ_LEN
+
+inputs["input_ids"] = torch.nn.functional.pad(inputs["input_ids"], (0, padded_len - input_ids_length), "constant", pad_token_id)
+inputs["attention_mask"] = torch.nn.functional.pad(inputs["attention_mask"], (0, padded_len - input_ids_length), "constant", 0)
+
+for key, value in inputs.items():
+ inputs[key] = np.array(value)
+
+vision_inputs = {
+ key: value
+ for key, value in inputs.items()
+ if key in {"pixel_values", "image_masks", "image_input_idx", "valid_idx", "aspect_ratio_ids", "aspect_ratio_mask"}
+}
+for key in {"pixel_values", "image_masks"}:
+ if key in vision_inputs:
+ vision_inputs[key] = vision_inputs[key].astype("float16")
+
+vision_outputs = {}
+if vision_inputs:
+ vision_outputs = vision_session.run(vision_inputs)
+
+lang_inputs = {key: value for key, value in inputs.items() if key not in vision_inputs}
+if "position_ids" not in lang_inputs:
+ lang_inputs["position_ids"] = np.where(
+ lang_inputs.pop("attention_mask"), np.arange(padded_len), -1
+ )
+else:
+ lang_inputs.pop("attention_mask", None)
+lang_inputs["image_idx"] = np.array([[0]])
+if not skip_vision:
+ lang_inputs["vision_embeds"] = vision_outputs["vision_embeds"]
+
+generation_len = 100 if not skip_vision else 512
+all_outputs = []
+lang_prefill_session.set_buffers(vision_outputs)
+chunk_inputs = lang_inputs.copy()
+for chunk_idx in range(num_chunks):
+ chunk_inputs["input_ids"] = lang_inputs["input_ids"][:, chunk_idx * PREFILL_SEQ_LEN : (chunk_idx + 1) * PREFILL_SEQ_LEN]
+ chunk_inputs["position_ids"] = lang_inputs["position_ids"][..., chunk_idx * PREFILL_SEQ_LEN : (chunk_idx + 1) * PREFILL_SEQ_LEN]
+ outputs = lang_prefill_session.run(chunk_inputs)
+ _update_retained_states(chunk_inputs, outputs)
+ chunk_inputs["image_idx"] = outputs["image_idx_output"]
+
+all_outputs.append(np.argmax(outputs["logits"]))
+decode_inputs = {
+ "input_ids": np.argmax(outputs["logits"]).reshape(BS, 1),
+ "position_ids": np.max(lang_inputs["position_ids"], axis=-1, keepdims=True) + 1,
+}
+_update_retained_states(decode_inputs, outputs)
+decode_inputs["image_idx"] = outputs["image_idx_output"]
+if not skip_vision:
+ decode_inputs["vision_embeds"] = outputs["vision_embeds_RetainedState"]
+decode_out = lang_decode_session.run(decode_inputs)
+all_outputs.append(np.argmax(decode_out["logits"]))
+position_ids = np.max(decode_inputs["position_ids"], axis=-1, keepdims=True) + 1
+loop_decode_inputs = {
+ "input_ids": np.argmax(decode_out["logits"]).reshape(BS, 1),
+ "position_ids": position_ids,
+}
+_update_retained_states(loop_decode_inputs, decode_out)
+loop_decode_inputs["image_idx"] = decode_out["image_idx_output"]
+if not skip_vision:
+ loop_decode_inputs["vision_embeds"] = decode_out["vision_embeds_RetainedState"]
+
+for _ in range(generation_len - 2):
+ decode_out = lang_decode_session.run(loop_decode_inputs)
+ all_outputs.append(np.argmax(decode_out["logits"]))
+ position_ids += 1
+ _update_retained_states(loop_decode_inputs, decode_out)
+ loop_decode_inputs.update(
+ {
+ "input_ids": np.argmax(decode_out["logits"]).reshape(BS, 1),
+ "position_ids": position_ids,
+ }
)
- streamer = TextStreamer(tokenizer)
- output = qeff_model.generate(inputs=inputs, generation_len=100, streamer=streamer)
- print(output.generated_ids)
- print(tokenizer.batch_decode(output.generated_ids))
- print(output)
+print(tokenizer.decode(np.asarray(all_outputs).reshape(-1).tolist()))
diff --git a/examples/image_text_to_text/models/qwen3_vl_moe/qwen3_vl_moe.py b/examples/image_text_to_text/models/qwen3_vl_moe/qwen3_vl_moe.py
index 649f51e6e5..b3b0339e5b 100644
--- a/examples/image_text_to_text/models/qwen3_vl_moe/qwen3_vl_moe.py
+++ b/examples/image_text_to_text/models/qwen3_vl_moe/qwen3_vl_moe.py
@@ -26,6 +26,7 @@
attn_implementation="eager",
kv_offload=True,
config=config,
+ weight_free=True,
# For CCL activation
# qaic_config={
# "ccl_enabled": True,
@@ -59,6 +60,7 @@
skip_vision=True,
mos=1,
use_onnx_subfunctions=True,
+ dynamo=True,
# comp_ctx_lengths_prefill=comp_ctx_lengths_prefill,
# comp_ctx_lengths_decode=comp_ctx_lengths_decode,
)
@@ -96,7 +98,7 @@
prefill_seq_len=128,
ctx_len=4096,
num_cores=16,
- num_devices=4,
+ num_devices=2,
height=354,
width=536,
split_model_io=True,
@@ -105,6 +107,7 @@
aic_enable_depth_first=True,
mos=1,
use_onnx_subfunctions=True,
+ dynamo=True,
# comp_ctx_lengths_prefill=comp_ctx_lengths_prefill,
# comp_ctx_lengths_decode=comp_ctx_lengths_decode,
)
diff --git a/examples/image_text_to_text/models/qwen3vl/qwen3_vl.py b/examples/image_text_to_text/models/qwen3vl/qwen3_vl.py
index 6aeb3efd6d..f7417957bc 100644
--- a/examples/image_text_to_text/models/qwen3vl/qwen3_vl.py
+++ b/examples/image_text_to_text/models/qwen3vl/qwen3_vl.py
@@ -5,139 +5,175 @@
#
# -----------------------------------------------------------------------------
+"""Dynamo-based export and inference for Qwen3-VL-MoE on Cloud AI 100.
+
+Requires PyTorch >= 2.13. Install dependencies before running:
+ pip install -r examples/dynamo/image_text_to_text/requirements.txt
+"""
+
+import argparse
+
import requests
import transformers
from PIL import Image
from qwen_vl_utils import process_vision_info
-from transformers import AutoConfig, AutoProcessor, TextStreamer
+from transformers import AutoConfig, AutoProcessor
from QEfficient import QEFFAutoModelForImageTextToText
+from QEfficient.utils import constants
-model_id = "Qwen/Qwen3-VL-32B-Instruct"
-config = AutoConfig.from_pretrained(model_id)
-
-# config.vision_config.depth = 9
-# config.text_config.num_hidden_layers = 1
-# config.vision_config.deepstack_visual_indexes = [8]
-
-qeff_model = QEFFAutoModelForImageTextToText.from_pretrained(
- model_id,
- attn_implementation="eager",
- kv_offload=True,
- config=config,
- # # For CCL activation
- # qaic_config={
- # "ccl_enabled": True,
- # },
-)
-tokenizer = transformers.AutoTokenizer.from_pretrained(model_id)
-processor = AutoProcessor.from_pretrained(model_id)
-### use skip_vision=Ture, if want to run only text, else false ###
-skip_vision = True
-
-# Compute-Context-Length (CCL) lists for prefill and decode. When both are None and
-# ccl_enabled=True, they are auto-generated from ctx_len.
-# comp_ctx_lengths_prefill = [2048]
-# comp_ctx_lengths_decode = [65536]
-
-if skip_vision:
- ## Only Text ##
-
- ## Set Batch_Size ##
- batch_size = 1
- qeff_model.compile(
- batch_size=batch_size,
- prefill_seq_len=128,
- ctx_len=4096,
- num_cores=16,
- num_devices=4,
- height=354,
- width=536,
- mxfp6_matmul=True,
- aic_enable_depth_first=True,
- skip_vision=True,
- mos=1,
- use_onnx_subfunctions=True,
- # comp_ctx_lengths_prefill=comp_ctx_lengths_prefill,
- # comp_ctx_lengths_decode=comp_ctx_lengths_decode,
- )
- messages = [
- {
- "role": "user",
- "content": [
- {"type": "text", "text": "Tell me about yourself."},
- ],
- },
- ]
-
- messages = [messages] * batch_size
-
- inputs = processor.apply_chat_template(
- messages,
- add_generation_prompt=True,
- tokenize=True,
- return_dict=True,
- return_tensors="pt",
- )
- inputs = qeff_model.model.prepare_inputs_for_generation(inputs=inputs, prefill_seq_len=128, batch_size=batch_size)
- streamer = TextStreamer(tokenizer)
- output = qeff_model.generate(inputs=inputs, generation_len=100)
- print(output.generated_ids)
- print(processor.tokenizer.batch_decode(output.generated_ids))
- print(output)
-
-else:
- batch_size = 1
- ## Vision + Text ##
- qeff_model.compile(
- batch_size=batch_size,
- prefill_seq_len=128,
- ctx_len=4096,
- num_cores=16,
- num_devices=4,
- height=354,
- width=536,
- split_model_io=True,
- mxfp6_matmul=True,
- mxint8_kv_cache=True,
- aic_enable_depth_first=True,
- mos=1,
- use_onnx_subfunctions=False,
- # comp_ctx_lengths_prefill=comp_ctx_lengths_prefill,
- # comp_ctx_lengths_decode=comp_ctx_lengths_decode,
- )
+def load_image(image_url: str, width: int, height: int) -> Image.Image:
+ """Load a remote image, falling back to a deterministic local image."""
+ try:
+ response = requests.get(image_url, stream=True, timeout=30)
+ response.raise_for_status()
+ return Image.open(response.raw).convert("RGB")
+ except requests.RequestException:
+ return Image.new("RGB", (width, height), color=(120, 70, 200))
- ### IMAGE + TEXT ###
- image_url = "https://picsum.photos/id/237/536/354"
- image = Image.open(requests.get(image_url, stream=True).raw)
+def maybe_reduce_config(config, *, reduce_layers: bool, vision_depth: int, text_layers: int):
+ """Apply the small bring-up config used by the Qwen3-VL-MoE example."""
+ if not reduce_layers:
+ return config
- messages_1 = [
- {
- "role": "user",
- "content": [
- {"type": "image", "image": image},
- {"type": "text", "text": "Descibe the image in details."},
- ],
- },
- ]
+ config.vision_config.depth = vision_depth
+ config.text_config.num_hidden_layers = text_layers
+ config.vision_config.deepstack_visual_indexes = [vision_depth - 1]
+ return config
- messages = [messages_1] * batch_size
- texts = [processor.apply_chat_template(msg, tokenize=False, add_generation_prompt=True) for msg in messages]
+def main():
+ parser = argparse.ArgumentParser(
+ description="Dynamo-based dual-QPC VLM export and inference for Qwen3-VL-MoE on Cloud AI 100.",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ parser.add_argument("--model-name", type=str, default="Qwen/Qwen3-VL-2B-Instruct")
+ parser.add_argument("--prompt", type=str, default="Describe all the colors seen in the image.")
+ parser.add_argument("--image-url", type=str, default="https://picsum.photos/id/237/536/354")
+ parser.add_argument("--height", type=int, default=354)
+ parser.add_argument("--width", type=int, default=536)
+ parser.add_argument("--batch-size", type=int, default=1)
+ parser.add_argument("--prefill-seq-len", type=int, default=128)
+ parser.add_argument("--ctx-len", type=int, default=4096)
+ parser.add_argument("--generation-len", type=int, default=100)
+ parser.add_argument("--num-cores", type=int, default=constants.DEFAULT_AIC_NUM_CORES)
+ parser.add_argument("--num-devices", type=int, default=4)
+ parser.add_argument("--mos", type=int, default=1)
+ parser.add_argument("--aic-hw-version", type=str, default=constants.DEFAULT_AIC_HW_VERSION)
+ parser.add_argument("--vision-depth", type=int, default=9)
+ parser.add_argument("--text-layers", type=int, default=1)
+ parser.add_argument(
+ "--reduce-layers",
+ action=argparse.BooleanOptionalAction,
+ default=False,
+ help="Use a reduced-layer config for faster bring-up.",
+ )
+ parser.add_argument(
+ "--weight-free",
+ action=argparse.BooleanOptionalAction,
+ default=True,
+ help="Build the model on meta tensors and load weights at compile time.",
+ )
+ parser.add_argument(
+ "--skip-vision",
+ action="store_true",
+ help="Compile and run the text path only.",
+ )
+ args = parser.parse_args()
+
+ config = AutoConfig.from_pretrained(args.model_name)
+ config = maybe_reduce_config(
+ config,
+ reduce_layers=args.reduce_layers,
+ vision_depth=args.vision_depth,
+ text_layers=args.text_layers,
+ )
- image_inputs, video_inputs = process_vision_info(messages)
- inputs = processor(
- text=texts,
- images=image_inputs,
- videos=video_inputs,
- padding=True,
- return_tensors="pt",
+ qeff_model = QEFFAutoModelForImageTextToText.from_pretrained(
+ args.model_name,
+ attn_implementation="eager",
+ kv_offload=True,
+ config=config,
+ weight_free=args.weight_free,
)
- inputs = qeff_model.model.prepare_inputs_for_generation(inputs=inputs, prefill_seq_len=128, batch_size=batch_size)
- streamer = TextStreamer(tokenizer)
- output = qeff_model.generate(inputs=inputs, generation_len=100)
+ tokenizer = transformers.AutoTokenizer.from_pretrained(args.model_name)
+ processor = AutoProcessor.from_pretrained(args.model_name)
+
+ compile_kwargs = {
+ "batch_size": args.batch_size,
+ "prefill_seq_len": args.prefill_seq_len,
+ "ctx_len": args.ctx_len,
+ "num_cores": args.num_cores,
+ "num_devices": args.num_devices,
+ "height": args.height,
+ "width": args.width,
+ "mxfp6_matmul": True,
+ "aic_enable_depth_first": True,
+ "mos": args.mos,
+ "use_onnx_subfunctions": True,
+ "aic_hw_version": args.aic_hw_version,
+ "dynamo": True,
+ }
+ if args.skip_vision:
+ compile_kwargs["skip_vision"] = True
+ else:
+ compile_kwargs.update({"split_model_io": True, "mxint8_kv_cache": True})
+
+ qpc_paths = qeff_model.compile(**compile_kwargs)
+ print(f"Model compiled to: {qpc_paths}")
+ if args.weight_free:
+ print(f"Weight specs: {qeff_model.weight_spec_path}")
+
+ if args.skip_vision:
+ messages = [
+ {
+ "role": "user",
+ "content": [{"type": "text", "text": args.prompt}],
+ }
+ ]
+ messages = [messages] * args.batch_size
+ inputs = processor.apply_chat_template(
+ messages,
+ add_generation_prompt=True,
+ tokenize=True,
+ return_dict=True,
+ return_tensors="pt",
+ )
+ else:
+ image = load_image(args.image_url, args.width, args.height)
+ messages = [
+ {
+ "role": "user",
+ "content": [
+ {"type": "image", "image": image},
+ {"type": "text", "text": args.prompt},
+ ],
+ }
+ ]
+ messages = [messages] * args.batch_size
+ texts = [processor.apply_chat_template(msg, tokenize=False, add_generation_prompt=True) for msg in messages]
+ image_inputs, video_inputs = process_vision_info(messages)
+ inputs = processor(
+ text=texts,
+ images=image_inputs,
+ videos=video_inputs,
+ padding=True,
+ return_tensors="pt",
+ )
+
+ inputs = qeff_model.model.prepare_inputs_for_generation(
+ inputs=inputs,
+ prefill_seq_len=args.prefill_seq_len,
+ batch_size=args.batch_size,
+ )
+ output = qeff_model.generate(inputs=inputs, generation_len=args.generation_len)
+
print(output.generated_ids)
- print(processor.tokenizer.batch_decode(output.generated_ids))
+ print(tokenizer.batch_decode(output.generated_ids))
print(output)
+
+
+if __name__ == "__main__":
+ main()
diff --git a/tests/base/test_onnx_transforms.py b/tests/base/test_onnx_transforms.py
index cfa309f5b4..7619dd98fd 100644
--- a/tests/base/test_onnx_transforms.py
+++ b/tests/base/test_onnx_transforms.py
@@ -11,7 +11,9 @@
from QEfficient.base.onnx_transforms import (
FP16ClipTransform,
OnnxTransformPipeline,
+ QualifyOnnxNodeNamesTransform,
RenameWsubNodesTransform,
+ RewriteSequenceSplitGetItemTransform,
SplitTensorsTransform,
)
@@ -114,6 +116,71 @@ def test_rename_wsub_nodes_transform():
assert not RenameWsubNodesTransform.apply(model)
+def test_qualify_onnx_node_names_transform():
+ function_node = onnx.helper.make_node("Identity", ["x"], ["y"], name="node_Identity_0")
+ function = onnx.helper.make_function(
+ "pkg.torch.__subgraph__",
+ "DecoderLayer",
+ ["x"],
+ ["y"],
+ [function_node],
+ [onnx.helper.make_opsetid("", 17)],
+ )
+ main_node = onnx.helper.make_node("Identity", ["input"], ["output"])
+ model = onnx.helper.make_model(
+ onnx.helper.make_graph([main_node], "test", [], []),
+ functions=[function],
+ opset_imports=[onnx.helper.make_opsetid("", 17)],
+ )
+
+ assert QualifyOnnxNodeNamesTransform.apply(model)
+ assert model.graph.node[0].name == "main/node_0_Identity"
+ assert model.functions[0].node[0].name == "DecoderLayer/node_Identity_0"
+ assert not QualifyOnnxNodeNamesTransform.apply(model)
+
+
+def test_rewrite_sequence_split_getitem_transform():
+ nodes = [
+ onnx.helper.make_node(
+ "Constant",
+ [],
+ ["split_size"],
+ name="split_size",
+ value=onnx.numpy_helper.from_array(np.asarray(4, dtype=np.int64)),
+ ),
+ onnx.helper.make_node("aten_split", ["x", "split_size"], ["parts"], name="split"),
+ onnx.helper.make_node(
+ "Constant",
+ [],
+ ["index"],
+ name="index",
+ value=onnx.numpy_helper.from_array(np.asarray(1, dtype=np.int64)),
+ ),
+ onnx.helper.make_node("aten_getitem", ["parts", "index"], ["y"], name="getitem"),
+ ]
+ function = onnx.helper.make_function(
+ "pkg.torch.__subgraph__",
+ "DecoderLayer",
+ ["x"],
+ ["y"],
+ nodes,
+ [onnx.helper.make_opsetid("", 17)],
+ )
+ model = onnx.helper.make_model(
+ onnx.helper.make_graph([], "test", [], []),
+ functions=[function],
+ opset_imports=[onnx.helper.make_opsetid("", 17)],
+ )
+
+ assert RewriteSequenceSplitGetItemTransform.apply(model)
+ rewritten = model.functions[0].node
+ assert rewritten[-1].op_type == "Slice"
+ assert [node.op_type for node in rewritten].count("Slice") == 1
+ assert rewritten[-1].input[0] == "x"
+ assert rewritten[-1].output[0] == "y"
+ assert not RewriteSequenceSplitGetItemTransform.apply(model)
+
+
def test_split_tensors_transform(tmp_path):
external_tensors_file = "tensors.raw"
test_onnx = onnx.parser.parse_model(f"""
diff --git a/tests/unit_test/models/test_model_quickcheck.py b/tests/unit_test/models/test_model_quickcheck.py
index 5fe9217794..7112c61af8 100644
--- a/tests/unit_test/models/test_model_quickcheck.py
+++ b/tests/unit_test/models/test_model_quickcheck.py
@@ -2341,6 +2341,55 @@ def test_qwen3_5_moe_get_submodules_for_export_keeps_decoder_layer_for_mixed_lay
assert wrapper.get_submodules_for_export() == {QEffQwen3_5MoeDecoderLayer}
+@pytest.mark.parametrize(
+ "model_module,model_class_name",
+ [
+ ("QEfficient.transformers.models.qwen3_5.modeling_qwen3_5", "QEffQwen3_5ForConditionalGeneration"),
+ (
+ "QEfficient.transformers.models.qwen3_5_moe.modeling_qwen3_5_moe",
+ "QEffQwen3_5MoeForConditionalGeneration",
+ ),
+ ],
+)
+def test_qwen3_5_decode_dummy_inputs_separate_seq_len_and_ctx_len(model_module, model_class_name):
+ """Decode examples keep one token while full-attention cache examples keep ctx_len."""
+ from importlib import import_module
+ from types import SimpleNamespace
+
+ model_class = getattr(import_module(model_module), model_class_name)
+ layer_types = ["linear_attention", "full_attention"]
+ text_config = SimpleNamespace(
+ num_hidden_layers=len(layer_types),
+ layer_types=layer_types,
+ num_key_value_heads=2,
+ num_attention_heads=8,
+ head_dim=256,
+ hidden_size=1024,
+ torch_dtype=torch.float32,
+ )
+ linear_attn = SimpleNamespace(
+ conv_dim=16,
+ conv_kernel_size=4,
+ num_k_heads=16,
+ num_v_heads=16,
+ head_k_dim=128,
+ head_v_dim=128,
+ )
+ layers = [SimpleNamespace(linear_attn=linear_attn) for _ in layer_types]
+ model = SimpleNamespace(
+ config=SimpleNamespace(text_config=text_config),
+ language_model=SimpleNamespace(layers=layers),
+ )
+ wrapper = SimpleNamespace(model=model)
+
+ inputs = model_class.get_dummy_inputs(wrapper, kv_offload=True, prefill_seq_len=1, ctx_len=4096)
+ lang_inputs = inputs["lang"]
+
+ assert lang_inputs["input_ids"].shape[-1] == 1
+ assert lang_inputs["position_ids"].shape[-1] == 1
+ assert lang_inputs["past_key_values"][1][0].shape[2] == 4096
+
+
def test_qwen3_5_moe_get_specializations_supports_multi_resolution():
from types import SimpleNamespace