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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions QEfficient/exporter/weight_free/checkpoint_key_resolver.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,8 @@
"original_inv_freq",
"embed_positions",
"embed_scale",
"position_ids",
"token_type_ids",
}


Expand Down
116 changes: 114 additions & 2 deletions QEfficient/transformers/models/bert/modeling_bert.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,12 +16,24 @@

Fix: override `_create_attention_masks` to use `_prepare_4d_attention_mask`
(standard tensor ops, fully ONNX-traceable) for the encoder (non-decoder) path.

Separately, `RobertaEmbeddings` / `XLMRobertaEmbeddings` / `NomicBertEmbeddings`
`.forward` build `token_type_ids` from a non-persistent `token_type_ids` buffer
(created on a concrete device) gathered against `position_ids`. Under
FakeTensor/meta-device tracing (e.g. dynamo + use_onnx_subfunctions,
weight-free export) `position_ids` can be a fake tensor on the `meta` device
while the buffer stays on a concrete device, so `torch.gather` fails with
"found at least two devices". The `QEff*Embeddings` classes below override
`forward` to move the buffer to `position_ids.device` before the gather;
otherwise identical to the upstream implementation.
"""

import torch
from transformers.modeling_attn_mask_utils import _prepare_4d_attention_mask
from transformers.models.bert.modeling_bert import BertModel
from transformers.models.roberta.modeling_roberta import RobertaModel
from transformers.models.xlm_roberta.modeling_xlm_roberta import XLMRobertaModel
from transformers.models.nomic_bert.modeling_nomic_bert import NomicBertEmbeddings
from transformers.models.roberta.modeling_roberta import RobertaEmbeddings, RobertaModel
from transformers.models.xlm_roberta.modeling_xlm_roberta import XLMRobertaEmbeddings, XLMRobertaModel


class _QEffBertFamilyMixin:
Expand Down Expand Up @@ -76,3 +88,103 @@ class QEffRobertaModel(_QEffBertFamilyMixin, RobertaModel):

class QEffXLMRobertaModel(_QEffBertFamilyMixin, XLMRobertaModel):
pass


class _QEffRobertaFamilyEmbeddingsMixin:
"""
Shared fixed `forward` for `RobertaEmbeddings` / `XLMRobertaEmbeddings` (identical
upstream implementations): moves the buffered `token_type_ids` to `position_ids.device`
before the `torch.gather` call. See module docstring for the FakeTensor/meta-device
rationale.
"""

def forward(
self,
input_ids=None,
token_type_ids=None,
position_ids=None,
inputs_embeds=None,
past_key_values_length=0,
):
if position_ids is None:
if input_ids is not None:
position_ids = self.create_position_ids_from_input_ids(
input_ids, self.padding_idx, past_key_values_length
)
else:
position_ids = self.create_position_ids_from_inputs_embeds(inputs_embeds, self.padding_idx)

if input_ids is not None:
input_shape = input_ids.size()
else:
input_shape = inputs_embeds.size()[:-1]

batch_size, seq_length = input_shape

if token_type_ids is None:
if hasattr(self, "token_type_ids"):
buffered_token_type_ids = self.token_type_ids.expand(position_ids.shape[0], -1)
buffered_token_type_ids = buffered_token_type_ids.to(position_ids.device)
buffered_token_type_ids = torch.gather(buffered_token_type_ids, dim=1, index=position_ids)
token_type_ids = buffered_token_type_ids.expand(batch_size, seq_length)
else:
token_type_ids = torch.zeros(input_shape, dtype=torch.long, device=self.position_ids.device)

if inputs_embeds is None:
inputs_embeds = self.word_embeddings(input_ids)
token_type_embeddings = self.token_type_embeddings(token_type_ids)
embeddings = inputs_embeds + token_type_embeddings

position_embeddings = self.position_embeddings(position_ids)
embeddings = embeddings + position_embeddings

embeddings = self.LayerNorm(embeddings)
embeddings = self.dropout(embeddings)
return embeddings


class QEffRobertaEmbeddings(_QEffRobertaFamilyEmbeddingsMixin, RobertaEmbeddings):
pass


class QEffXLMRobertaEmbeddings(_QEffRobertaFamilyEmbeddingsMixin, XLMRobertaEmbeddings):
pass


class QEffNomicBertEmbeddings(NomicBertEmbeddings):
"""
Fixes the same buffered-`token_type_ids` device mismatch as `QEffRobertaEmbeddings`,
for `NomicBertEmbeddings` (native in Transformers v5.5, generated from
`modular_nomic_bert.py`). Otherwise identical to the upstream implementation.
"""

def forward(
self,
input_ids=None,
token_type_ids=None,
position_ids=None,
inputs_embeds=None,
):
embeddings = inputs_embeds
if inputs_embeds is None:
embeddings = self.word_embeddings(input_ids)

input_shape = embeddings.shape[:-1]
device = embeddings.device

if token_type_ids is None:
if hasattr(self, "token_type_ids"):
buffered_token_type_ids = self.token_type_ids.expand(position_ids.shape[0], -1)
buffered_token_type_ids = buffered_token_type_ids.to(position_ids.device)
buffered_token_type_ids = torch.gather(buffered_token_type_ids, dim=1, index=position_ids)
token_type_ids = buffered_token_type_ids.expand(*input_shape)
else:
token_type_ids = torch.zeros(input_shape, dtype=torch.long, device=device)

token_type_embeddings = self.token_type_embeddings(token_type_ids)

embeddings = embeddings + token_type_embeddings
embeddings = self.LayerNorm(embeddings)
embeddings = self.dropout(embeddings)

return embeddings
38 changes: 31 additions & 7 deletions QEfficient/transformers/models/modeling_auto.py
Original file line number Diff line number Diff line change
Expand Up @@ -173,7 +173,7 @@ def _disable_unsupported_weight_free(kwargs: dict, qeff_auto_class_name: str) ->
return

logger.warning(
"weight_free=True is only supported for QEFFAutoModelForCausalLM; disabling it for %s.",
"weight_free=True is only supported for QEFFAutoModelForCausalLM, QEffAutoModel; disabling it for %s.",
qeff_auto_class_name,
)

Expand Down Expand Up @@ -394,7 +394,8 @@ 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 self.__class__ is not QEFFAutoModel:
_disable_unsupported_weight_free(kwargs, self.__class__.__name__)

if (
hasattr(model, "config")
Expand Down Expand Up @@ -538,6 +539,7 @@ class QEFFAutoModel(QEFFTransformersBase):
_pytorch_transforms = [CustomOpsTransform, AwqToMatmulNbitsTransform, GPTQToMatmulNbitsTransform]
# FP16Clip inlines external weights; without Split the saved protobuf exceeds 2GB for large embedders.
_onnx_transforms = [FP16ClipTransform, SplitTensorsTransform]
_checkpoint_transforms = [DtypeConversionCheckpointTransform]

def __init__(self, model: nn.Module, pooling=None, **kwargs):
"""
Expand Down Expand Up @@ -570,7 +572,7 @@ def __init__(self, model: nn.Module, pooling=None, **kwargs):

@classmethod
@with_replaced_quantizers
def from_pretrained(cls, pretrained_model_name_or_path, pooling=None, *args, **kwargs):
def from_pretrained(cls, pretrained_model_name_or_path, pooling=None, weight_free=False, *args, **kwargs):
"""
Load a QEfficient transformer model from a pretrained HuggingFace model or local path.

Expand All @@ -589,6 +591,8 @@ def from_pretrained(cls, pretrained_model_name_or_path, pooling=None, *args, **k
- "avg": Average pooling
- Callable: A custom pooling function
- None: No pooling applied. Default is None.
weight_free : bool, optional
If True, the model will be loaded in weight-free mode, which avoids materializing checkpoint weights. Default is False.
*args :
Positional arguments passed directly to `cls._hf_auto_class.from_pretrained`.
**kwargs :
Expand All @@ -602,9 +606,11 @@ 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 weight_free:
validate_dynamo_export_requirements("weight_free=True")

if kwargs.get("attn_implementation", None) not in {None, "eager"}:
logger.warning('Updating attn_implementation="eager"')

Expand All @@ -614,7 +620,13 @@ def from_pretrained(cls, pretrained_model_name_or_path, pooling=None, *args, **k
kwargs.update({"attn_implementation": "eager", "low_cpu_mem_usage": False})

_resolve_torch_dtype(kwargs)
model = cls._hf_auto_class.from_pretrained(pretrained_model_name_or_path, *args, **kwargs)
if weight_free:
# Weight-free mode: build the model on the meta device so no
# checkpoint weights are ever materialized here. The real weights
# are supplied later at export time via pretrained_model_name_or_path.
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, *args, **kwargs)

# This is support models that should be classified to in a different auto class but transformers load them via this class
kv_offload = kwargs.pop("kv_offload", None)
Expand All @@ -626,7 +638,13 @@ def from_pretrained(cls, pretrained_model_name_or_path, pooling=None, *args, **k
model, kv_offload=kv_offload, **kwargs
)

return cls(model, pretrained_model_name_or_path=pretrained_model_name_or_path, pooling=pooling, **kwargs)
return cls(
model,
pretrained_model_name_or_path=pretrained_model_name_or_path,
pooling=pooling,
weight_free=weight_free,
**kwargs,
)

@property
def get_model_config(self) -> dict:
Expand All @@ -640,7 +658,7 @@ def get_model_config(self) -> dict:
"""
return self.model.config.__dict__

def export(self, export_dir: str | None = None, **kwargs) -> str:
def export(self, export_dir: str | None = None, dynamo: bool = False, **kwargs) -> str:
"""
Export the model to ONNX format using ``torch.onnx.export``.

Expand All @@ -663,6 +681,11 @@ def export(self, export_dir: str | None = None, **kwargs) -> str:
bs = constants.ONNX_EXPORT_EXAMPLE_BATCH_SIZE
seq_len = constants.ONNX_EXPORT_EXAMPLE_SEQ_LEN

dynamo = kwargs.get("dynamo", False) or self._weight_free
if dynamo:
# torch.export requires example inputs to satisfy dynamic_shapes min=2; gpt_oss non-CB keeps bs=1.
bs = max(2, bs)

example_inputs = {
"input_ids": torch.zeros((bs, seq_len), dtype=torch.int64),
"attention_mask": torch.ones((bs, seq_len), dtype=torch.int64),
Expand All @@ -677,6 +700,7 @@ def export(self, export_dir: str | None = None, **kwargs) -> str:
output_names=output_names,
dynamic_axes=dynamic_axes,
export_dir=export_dir,
dynamo=dynamo,
use_onnx_subfunctions=kwargs.get("use_onnx_subfunctions", False),
)

Expand Down
13 changes: 11 additions & 2 deletions QEfficient/transformers/models/pytorch_transforms.py
Original file line number Diff line number Diff line change
Expand Up @@ -216,6 +216,7 @@
except ImportError:
from transformers.models.qwen2_5_vl.modeling_qwen2_5_vl import Qwen2_5_VLRMSNorm as Qwen2_5RMSNorm
from transformers.models.bert.modeling_bert import BertModel
from transformers.models.nomic_bert.modeling_nomic_bert import NomicBertEmbeddings
from transformers.models.qwen3.modeling_qwen3 import (
Qwen3Attention,
Qwen3DecoderLayer,
Expand Down Expand Up @@ -288,7 +289,7 @@
Qwen3VLMoeVisionAttention,
Qwen3VLMoeVisionModel,
)
from transformers.models.roberta.modeling_roberta import RobertaModel
from transformers.models.roberta.modeling_roberta import RobertaEmbeddings, RobertaModel
from transformers.models.starcoder2.modeling_starcoder2 import (
Starcoder2Attention,
Starcoder2DecoderLayer,
Expand All @@ -313,7 +314,7 @@
WhisperModel,
WhisperPositionalEmbedding,
)
from transformers.models.xlm_roberta.modeling_xlm_roberta import XLMRobertaModel
from transformers.models.xlm_roberta.modeling_xlm_roberta import XLMRobertaEmbeddings, XLMRobertaModel

from QEfficient.base.pytorch_transforms import (
ExternalModuleMapperTransform,
Expand All @@ -325,7 +326,10 @@
from QEfficient.transformers.embeddings.embedding_utils import POOLING_MAP, PooledModel, validate_user_pooling_function
from QEfficient.transformers.models.bert.modeling_bert import (
QEffBertModel,
QEffNomicBertEmbeddings,
QEffRobertaEmbeddings,
QEffRobertaModel,
QEffXLMRobertaEmbeddings,
QEffXLMRobertaModel,
)
from QEfficient.transformers.models.codegen.modeling_codegen import (
Expand Down Expand Up @@ -724,6 +728,11 @@ class CustomOpsTransform(ModuleMappingTransform):
BertModel: QEffBertModel,
RobertaModel: QEffRobertaModel,
XLMRobertaModel: QEffXLMRobertaModel,
# *Embeddings: fix a FakeTensor/meta-device mismatch in the buffered
# token_type_ids gather (see QEff*Embeddings docstrings in modeling_bert.py).
RobertaEmbeddings: QEffRobertaEmbeddings,
XLMRobertaEmbeddings: QEffXLMRobertaEmbeddings,
NomicBertEmbeddings: QEffNomicBertEmbeddings,
Qwen3_5RMSNorm: GemmaCustomRMSNormAIC,
Qwen3_5MoeRMSNorm: GemmaCustomRMSNormAIC,
Qwen3_5RMSNormGated: QEffQwen3_5GatedDeltaNetCustomRMSNormAIC,
Expand Down
23 changes: 19 additions & 4 deletions examples/embeddings/text_embeddings.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,16 @@ def main():
default="32,64",
help="Sequence length(s) - single int (e.g., '32') or comma-separated list (e.g., '32,64')",
)
parser.add_argument(
"--weight-free",
action="store_true",
help="Build the model on meta tensors and load weights at compile time",
)
parser.add_argument(
"--use-onnx-subfunctions",
action="store_true",
help="Use subfunctions while exporting",
)
args = parser.parse_args()

# Parse seq_len argument
Expand All @@ -67,15 +77,20 @@ def main():
# You can specify the pooling strategy either as a string (e.g., "max") or by passing a custom pooling function.
# If no pooling is specified, the model will return its default output (typically token embeddings).
if args.pooling == "max":
qeff_model = AutoModel.from_pretrained(args.model_name, pooling=max_pooling)
qeff_model = AutoModel.from_pretrained(args.model_name, pooling=max_pooling, weight_free=args.weight_free)
elif args.pooling == "mean":
qeff_model = AutoModel.from_pretrained(args.model_name, pooling="mean")
qeff_model = AutoModel.from_pretrained(args.model_name, pooling="mean", weight_free=args.weight_free)
else:
qeff_model = AutoModel.from_pretrained(args.model_name)
qeff_model = AutoModel.from_pretrained(args.model_name, weight_free=args.weight_free)

# Compile the model
# seq_len can be a list of seq_len or single int
qeff_model.compile(num_cores=args.num_cores, seq_len=seq_len)
qeff_model.compile(
num_cores=args.num_cores,
seq_len=seq_len,
dynamo=args.weight_free,
use_onnx_subfunctions=args.use_onnx_subfunctions,
)

# Tokenize sentences
encoded_input = tokenizer(args.sentences, return_tensors="pt")
Expand Down
Loading
Loading