Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
107 commits
Select commit Hold shift + click to select a range
787b2ae
Upgrade transformers 5.3.0 dependency
longlee0622 Apr 8, 2026
4a002b1
fix(auto_deploy): fallback SDPA mask patch when executorch helper rem…
longlee0622 Apr 8, 2026
cf162a5
fix: import AutoModelForImageTextToText for Transformers v5
longlee0622 Apr 8, 2026
8797b84
fix: shim get_parameter_device/dtype for Transformers v5
longlee0622 Apr 8, 2026
5236fc0
fix: pass exist_ok to AutoConfig.register for Transformers 5.3
longlee0622 Apr 9, 2026
6e92c3e
fix: replace removed load_sharded_checkpoint for Transformers 5.3
longlee0622 Apr 9, 2026
3e08573
format code
longlee0622 Apr 9, 2026
3540592
test: replace HybridCache with StaticCache helper for Transformers v5
longlee0622 Apr 9, 2026
08ef431
fix: map Transformers v5 rope_scaling type default to RoPE none
longlee0622 Apr 9, 2026
b7c13cd
fix: read HF rope_theta from rope_parameters for Transformers v5
longlee0622 Apr 9, 2026
3f931e6
test: compat helpers for DynamicCache legacy API (Transformers v5)
longlee0622 Apr 9, 2026
7009978
fix(auto_deploy): patch BambaModel when _update_causal_mask is absent
longlee0622 Apr 9, 2026
27feaac
fix: add SlidingWindowCache compatibility shim for Transformers v5
longlee0622 Apr 9, 2026
9026ef3
test: update test_gpt_attention rope config for Transformers v5
longlee0622 Apr 9, 2026
e7e0fc6
fix: map rope_type "default" to rope_gpt_neox in PositionEmbeddingType
longlee0622 Apr 9, 2026
df92b1d
test: fix test_gpt_attention_IFB for Transformers v5
longlee0622 Apr 9, 2026
6c65dc9
fix: use getattr for pad_token_id in MllamaConfig for Transformers v5
longlee0622 Apr 9, 2026
885d915
style: apply pre-commit formatting fixes
longlee0622 Apr 9, 2026
c430462
fix: prevent duplicate 'disable' kwarg in DisabledTqdm
longlee0622 Apr 10, 2026
1275756
fix: handle rope_scaling key changes in Qwen models for Transformers v5
longlee0622 Apr 10, 2026
ffad29a
fix: use getattr for tie_word_embeddings for Transformers v5
longlee0622 Apr 10, 2026
8277f32
fix: support per-layer-type RoPE config (Gemma3) for Transformers v5
longlee0622 Apr 10, 2026
77774a6
test: fix test_gpt_attention HF attention calls for Transformers v5
longlee0622 Apr 10, 2026
9c1c766
fix: pass return_dict=False in apply_chat_template for Transformers v5
longlee0622 Apr 10, 2026
6cf77e6
fix: pass return_dict=False in chat example for Transformers v5
longlee0622 Apr 10, 2026
63002af
fix: add return_dict=False to remaining apply_chat_template calls
longlee0622 Apr 10, 2026
88946ef
fix: replace batch_encode_plus removed in Transformers v5
longlee0622 Apr 10, 2026
aae7a8b
test: add rope_type key for scaled RoPE configs in attention tests
longlee0622 Apr 10, 2026
d21e6e4
test: use boi_token_index instead of bos_token_id in recursive config…
longlee0622 Apr 11, 2026
250f3ef
fix: ensure rope_scaling has "type" key for HF custom model code
longlee0622 Apr 11, 2026
2ddad18
fix: patch find_packed_sequence_indices for meta tensor compat in export
longlee0622 Apr 11, 2026
447e135
fix: patch torch.is_autocast_enabled for export with unknown device t…
longlee0622 Apr 11, 2026
c93e968
fix: remap HF fused expert weights in KimiK2 equivalence tests
longlee0622 Apr 11, 2026
9413f84
fix: handle NotImplementedError from tokenizer.vocab_size in transfor…
longlee0622 Apr 11, 2026
c02fa1b
fix: handle tuple router_logits from Qwen3 MoE gate in transformers 5.x
longlee0622 Apr 11, 2026
7c772b1
fix: initialize rope_scaling dict before setting type for Qwen VL models
longlee0622 Apr 11, 2026
095b220
fix: override _check_and_adjust_experts_implementation for Qwen3VL
longlee0622 Apr 11, 2026
ad8eaaf
test: handle no_init_weights removal in Transformers v5
longlee0622 Apr 11, 2026
ddd76aa
fix: unfuse HF transformers 5.x fused MoE weights for Mixtral
longlee0622 Apr 11, 2026
c36f9a5
fix: unfuse HF transformers 5.x fused MoE weights for Qwen MoE models
longlee0622 Apr 11, 2026
4228f60
fix: handle "default" rope_type in Phi-3 custom code by clearing rope…
longlee0622 Apr 11, 2026
e9c849a
fix: fall back to len(tokenizer) in get_vocab_size when vocab_size ra…
longlee0622 Apr 11, 2026
8188b99
fix: accept extra args in Qwen3VL _check_and_adjust_experts_implement…
longlee0622 Apr 11, 2026
7f3b124
fix: guard against non-iterable fused MixtralExperts in MoE patch
longlee0622 Apr 11, 2026
3c54588
fix: handle renamed top_k and fused experts in Qwen3Next MoE patch
longlee0622 Apr 11, 2026
d9ec20c
fix: apply DeepSeek V3 yarn RoPE override unconditionally
longlee0622 Apr 11, 2026
5cce714
fix: also apply rope_scaling fix to text_config for VL models
longlee0622 Apr 11, 2026
fddd714
fix: only clear rope_scaling for Phi-3, not all models
longlee0622 Apr 12, 2026
d7d8a55
fix: unfuse MoE weights and remap names in Qwen3VL MoE weight mapper
longlee0622 Apr 12, 2026
a333de0
fix: treat rope_type "default" as no scaling in config flattening
longlee0622 Apr 12, 2026
57b584e
fix: handle fused expert keys without .weight suffix in KimiK2 test
longlee0622 Apr 12, 2026
f10d94c
test: remove mm_token_type_ids from HF inputs for multimodal tests
longlee0622 Apr 12, 2026
4928996
fix: use getattr for rope_scaling in Qwen2VL model init
longlee0622 Apr 12, 2026
cb0a934
fix: handle tuple router_logits and renamed top_k in Mixtral/Qwen3 Mo…
longlee0622 Apr 12, 2026
b5d8d1f
fix: add fallback for missing token IDs in T5 enc-dec checkpoint conv…
longlee0622 Apr 12, 2026
8b2f3af
fix: safely access config in Qwen3Next MoE patch top_k fallback
longlee0622 Apr 12, 2026
4ee5c98
fix: only remove mm_token_type_ids for video modality in multimodal t…
longlee0622 Apr 12, 2026
7c3f6e0
fix: unfuse FP8 scale tensors for MoE experts in weight mapper
longlee0622 Apr 12, 2026
e7eb886
fix: clear default rope_scaling after AutoConfig.from_pretrained
longlee0622 Apr 12, 2026
877fe9f
fix: use gate_up_proj split for fused Qwen3Next experts in MoE patch
longlee0622 Apr 12, 2026
ff35efb
fix: use gate_up_proj split and safe config access for fused Mixtral …
longlee0622 Apr 12, 2026
c602017
fix: catch NotImplementedError from len(tokenizer) in get_vocab_size
longlee0622 Apr 12, 2026
73a997f
test: relax custom_role assertion for Transformers v5 chat templates
longlee0622 Apr 12, 2026
405d1ff
fix: use text_config for RoPE params in Qwen2.5-VL init_mrope_embedding
longlee0622 Apr 12, 2026
20efc01
test: encode prompts individually in Eagle3 test for Transformers v5
longlee0622 Apr 12, 2026
46b0bb2
multimodal fixes
2ez4bz Apr 16, 2026
71f8012
[None] Disable two-model eagle3 testing in CI
ziyixiong-nv Apr 17, 2026
f19d85d
fix: reload ByteLevel tokenizer via PreTrainedTokenizerFast for Llama…
dc3671 Apr 17, 2026
7bb04d9
fix: pass explicit head_dim for Qwen VL vision attention
longlee0622 Apr 19, 2026
b47efcf
fix: bump mistral-common to 1.9.1 for transformers 5.3.0 compatibility
longlee0622 Apr 19, 2026
f8fd44a
fix: populate missing rope_parameters for PixtralRotaryEmbedding
longlee0622 Apr 20, 2026
4f9357e
fix: use rope_config consistently in Qwen2VL init_mrope_embedding
longlee0622 Apr 20, 2026
eb43d33
[None][fix] Resolve Gemma 3 RoPE fields from nested rope_parameters
eopXD Apr 20, 2026
a572ced
autodeploy test fixes for transformers 5.3 upgrade
lucaslie Apr 21, 2026
a149063
format
longlee0622 Apr 21, 2026
b85bea6
[None][chore] Unwaive disagg tests for transformers upgrade
brb-nv Apr 21, 2026
da0d63d
fix: use text_config for Qwen2.5-VL inner LM config
longlee0622 Apr 22, 2026
b9b5d43
[None][fix] pass fused MoE weights through in Qwen3VL MoE weight mapper
longlee0622 Apr 22, 2026
d7ed9f4
[None][fix] transpose Qwen3VL MoE expert weights in weight mapper
longlee0622 Apr 22, 2026
06954c0
[None][fix] resolve Qwen2.5-VL vision RMSNorm eps from text_config
longlee0622 Apr 22, 2026
5163ced
Update tensorrt_llm/_torch/auto_deploy/export/library/transformers_ca…
longlee0622 Apr 22, 2026
3f15ab0
[None][fix] resolve Qwen2.5-VL vision attention max_position_embeddin…
longlee0622 Apr 23, 2026
8035495
[None][fix] fix LlavaNext dtype mismatch with transformers 5.x
longlee0622 Apr 23, 2026
c8e23d6
[None][fix] fix Qwen3VL video MRope position IDs in multimodal unit test
longlee0622 Apr 23, 2026
257aa26
[None][fix] fix Qwen2.5-VL and Qwen3-VL multimodal unit tests for tra…
longlee0622 Apr 23, 2026
63eb6d6
[None][fix] fix VilaMultimodalProjector for transformers 5.x compatib…
longlee0622 Apr 23, 2026
a019147
[None][fix] apply ByteLevel tokenizer fix in kv_cache_aware router
longlee0622 Apr 25, 2026
d491c4b
[None][chore] remove transformers v4 fallbacks from v5 upgrade
longlee0622 Apr 25, 2026
d6183a3
[None][chore] address review comments on transformers v5 upgrade
longlee0622 Apr 26, 2026
d3b3f8a
Update ATTRIBUTIONS
longlee0622 Apr 27, 2026
cf7710a
[None][chore] use upstream transformers configs for nemotron_h, exaon…
longlee0622 Apr 27, 2026
9e7530b
[None][chore] bump triton_backend transformers to 5.3.0
longlee0622 Apr 28, 2026
54adbff
[None][fix] fix exaone_moe weight mapper for transformers >=5.3 fused…
longlee0622 Apr 28, 2026
b21908a
[None][fix] fix T5Tokenizer sp_model removal in transformers >=5.3
longlee0622 Apr 28, 2026
a60db66
[None][fix] pass weight_mapper in exaone_moe HF accuracy test
longlee0622 Apr 28, 2026
e9d0edf
[None][perf] cache all_special_ids in convert_ids_to_tokens for slow …
longlee0622 Apr 30, 2026
9cd9854
[None][fix] read Llama4 RoPE rope_type from rope_parameters in transf…
longlee0622 May 6, 2026
9b3de5b
[None][fix] update Qwen3MoE / Starcoder2 custom models for transforme…
longlee0622 May 6, 2026
ba299d3
[None][fix] read rope_theta from rope_parameters in Qwen3MoE / Starco…
longlee0622 May 6, 2026
b56c6ea
[None][fix] strip processor-output keys from Qwen2/3-VL processor kwargs
longlee0622 May 6, 2026
8d03c9f
[None][fix] inject pad_token_id when loading Qwen2.5-Math-PRM-7B refe…
longlee0622 May 6, 2026
5d9f01a
[None][fix] unfuse HF Qwen3MoE expert params in NVFP4 IPC test helper
longlee0622 May 6, 2026
f9aa48f
[None][fix] bypass merged_typed_dict validation for Qwen2/3-VL processor
longlee0622 May 7, 2026
453ae2e
[None][fix] use dict form for Starcoder2 _tied_weights_keys (transfor…
longlee0622 May 7, 2026
9d6a629
[None][fix] init standalone HF Qwen3MoE block params in equivalence t…
longlee0622 May 7, 2026
8434568
[None][fix] re-init vendored RoPE buffers after loading Qwen2.5-Math-…
longlee0622 May 7, 2026
c9f01c0
[None][test] move LoRA e2e tests from A10 to A100 to avoid OOM
longlee0622 May 7, 2026
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion ATTRIBUTIONS-Python.md
Original file line number Diff line number Diff line change
Expand Up @@ -62159,7 +62159,7 @@ SOFTWARE.
- `Homepage`: https://github.com/EleutherAI/tqdm-multiprocess


## transformers (4.56.0)
## transformers (5.3.0)

### Licenses
License: `Apache 2.0 License`
Expand Down
3 changes: 2 additions & 1 deletion examples/apps/chat.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,7 +37,8 @@ def runsource(self,
self.history.append(message)

input = self.tokenizer.apply_chat_template(self.history,
add_generation_prompt=True)
add_generation_prompt=True,
return_dict=False)

output = self.llm.generate([input],
sampling_params=self.sampling_params)[0]
Expand Down
3 changes: 2 additions & 1 deletion examples/eagle/convert_checkpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@

import tensorrt_llm
from tensorrt_llm._deprecation import emit_engine_arch_deprecation
from tensorrt_llm._utils import get_hf_rope_theta
from tensorrt_llm.mapping import Mapping
from tensorrt_llm.models.eagle.config import EagleConfig
from tensorrt_llm.models.eagle.model import EagleForCausalLM
Expand Down Expand Up @@ -293,7 +294,7 @@ def copy(tensors):
args.rms_norm_eps = hf_config.rms_norm_eps
args.vocab_size = hf_config.vocab_size
args.rotary_scaling = hf_config.rope_scaling
args.rotary_base = hf_config.rope_theta
args.rotary_base = get_hf_rope_theta(hf_config, 10000.0)
args.n_positions = hf_config.max_position_embeddings
args.dtype = str(
hf_config.torch_dtype)[6:] if args.dtype == 'auto' else args.dtype
Expand Down
4 changes: 2 additions & 2 deletions examples/medusa/convert_checkpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@

import tensorrt_llm
from tensorrt_llm._deprecation import emit_engine_arch_deprecation
from tensorrt_llm._utils import numpy_to_torch
from tensorrt_llm._utils import get_hf_rope_theta, numpy_to_torch
from tensorrt_llm.logger import logger
from tensorrt_llm.mapping import Mapping
from tensorrt_llm.models import (LLaMAForCausalLM, PretrainedConfig,
Expand Down Expand Up @@ -209,7 +209,7 @@ def main():
args.rms_norm_eps = hf_config.rms_norm_eps
args.vocab_size = hf_config.vocab_size
args.n_positions = hf_config.max_position_embeddings
args.rotary_base = hf_config.rope_theta
args.rotary_base = get_hf_rope_theta(hf_config, 10000.0)
args.rotary_scaling = hf_config.rope_scaling

elif args.meta_ckpt_dir is not None:
Expand Down
4 changes: 2 additions & 2 deletions examples/models/contrib/dbrx/convert_checkpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,7 @@

import tensorrt_llm
from tensorrt_llm._deprecation import emit_engine_arch_deprecation
from tensorrt_llm._utils import release_gc
from tensorrt_llm._utils import get_hf_rope_theta, release_gc
from tensorrt_llm.layers import MoeConfig
from tensorrt_llm.mapping import Mapping
from tensorrt_llm.models.convert_utils import (generate_int8,
Expand Down Expand Up @@ -557,7 +557,7 @@ def execute(workers, func, hf_model):
args.moe_top_k = 1
args.clip_qkv = hf_config.attn_config.clip_qkv
args.hidden_act = 'swiglu'
args.rotary_base = hf_config.attn_config.rope_theta
args.rotary_base = get_hf_rope_theta(hf_config.attn_config, 10000.0)
args.moe_config = MoeConfig(
num_experts=args.moe_num_experts,
top_k=args.moe_top_k,
Expand Down
16 changes: 9 additions & 7 deletions examples/models/core/enc_dec/convert_checkpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -143,15 +143,17 @@ def parse_t5_config_by_component(config, component, args):
component_config.encoder_head_size = config.getint(
'encoder', 'd_kv')
component_config.decoder_start_token_id = config.getint(
'decoder', 'decoder_start_token_id')
component_config.eos_token_id = config.getint(
'decoder', 'eos_token_id')
bos_token_id = config.get('decoder', 'bos_token_id')
'decoder', 'decoder_start_token_id', fallback=0)
component_config.eos_token_id = config.getint('decoder',
'eos_token_id',
fallback=1)
bos_token_id = config.get('decoder', 'bos_token_id', fallback=None)
# T5 does not have bos_token_id
component_config.bos_token_id = int(
bos_token_id) if bos_token_id != "None" else None
component_config.pad_token_id = config.getint(
'decoder', 'pad_token_id')
bos_token_id) if bos_token_id not in (None, "None") else None
component_config.pad_token_id = config.getint('decoder',
'pad_token_id',
fallback=0)

else:
assert False, 'Unsupported component!'
Expand Down
4 changes: 2 additions & 2 deletions examples/models/core/internlm2/convert_checkpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,7 @@

import tensorrt_llm
from tensorrt_llm._deprecation import emit_engine_arch_deprecation
from tensorrt_llm._utils import release_gc
from tensorrt_llm._utils import get_hf_rope_theta, release_gc
from tensorrt_llm.mapping import Mapping
from tensorrt_llm.models.llama import convert

Expand Down Expand Up @@ -480,7 +480,7 @@ def convert_from_hf(hf_model,
'norm_epsilon': hf_config.rms_norm_eps,
'vocab_size': hf_config.vocab_size,
'position_embedding_type': 'rope_gpt_neox',
'rotary_base': hf_config.rope_theta,
'rotary_base': get_hf_rope_theta(hf_config, 10000.0),
'max_position_embeddings': hf_config.max_position_embeddings,
'hidden_act': hf_config.hidden_act,
'use_parallel_embedding': args.use_parallel_embedding,
Expand Down
4 changes: 2 additions & 2 deletions requirements.txt
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,7 @@ nvidia-modelopt[torch]~=0.37.0
# torch 2.10.0+cu130 depends on nvidia-nccl-cu13==2.28.9
nvidia-nccl-cu13>=2.28.9,<=2.29.2
nvidia-cuda-nvrtc
transformers==4.57.3
transformers==5.3.0
prometheus_client
prometheus_fastapi_instrumentator
pydantic>=2.9.1
Expand Down Expand Up @@ -77,7 +77,7 @@ numexpr
partial_json_parser
apache-tvm-ffi==0.1.6 # used for reduce nvidia-cutlass-dsl host overhead
torch-c-dlpack-ext==0.1.3 # used for reduce nvidia-cutlass-dsl host overhead, optional package for improved torch tensor calling perf
mistral-common==1.8.6
mistral-common==1.9.1
torchao>=0.14.1,<0.16.0
cuda-core
llist
Expand Down
29 changes: 22 additions & 7 deletions tensorrt_llm/_torch/attention_backend/interface.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,9 +14,10 @@
from ..speculative.interface import SpecMetadata
from ..speculative.spec_tree_manager import SpecTreeManager

from tensorrt_llm._utils import maybe_pin_memory
from tensorrt_llm._utils import get_hf_rope_theta, maybe_pin_memory
from tensorrt_llm.functional import (PositionEmbeddingType, RopeEmbeddingUtils,
RotaryScalingType)
from tensorrt_llm.logger import logger
from tensorrt_llm.mapping import Mapping
from tensorrt_llm.models.modeling_utils import QuantConfig

Expand Down Expand Up @@ -483,10 +484,23 @@ def from_config(config) -> "RopeParams":

hf_rope_parameters = getattr(config, 'rope_parameters', None)
if hf_rope_parameters is not None:
assert not set(hf_rope_parameters.keys()).issubset(
ALLOWED_ATTENTION_LAYER_TYPES), (
"Per-layer-type RoPE configuration is not supported yet.")
config.update(hf_rope_parameters)
if set(hf_rope_parameters.keys()).issubset(
ALLOWED_ATTENTION_LAYER_TYPES):
# Per-layer-type RoPE config (e.g. Gemma3 in transformers 5.x).
# Pick "full_attention" as the default; callers override theta
# for sliding-window layers independently.
if "full_attention" in hf_rope_parameters:
flat = hf_rope_parameters["full_attention"]
else:
fallback_key = next(iter(hf_rope_parameters))
logger.warning(
f"Per-layer-type rope_parameters has no 'full_attention' entry; "
f"falling back to '{fallback_key}'. Available layer types: "
f"{list(hf_rope_parameters.keys())}.")
flat = hf_rope_parameters[fallback_key]
config.update(flat)
else:
config.update(hf_rope_parameters)

# get rotary parameters.
hidden_size = config.hidden_size
Expand All @@ -496,7 +510,7 @@ def from_config(config) -> "RopeParams":
head_dim = hidden_size // num_attention_heads
rope_scaling = getattr(config, 'rope_scaling', None)
rope_params.max_positions = config.max_position_embeddings
rope_params.theta = getattr(config, 'rope_theta', 10000.0)
rope_params.theta = get_hf_rope_theta(config, 10000.0)
rope_percentage = (getattr(config, 'rotary_pct', None)
or getattr(config, 'partial_rotary_factor', None)
or 1.0)
Expand Down Expand Up @@ -534,7 +548,8 @@ def from_config(config) -> "RopeParams":
rope_params.short_factor = tuple(rope_scaling["short_factor"])
if "long_factor" in rope_scaling:
rope_params.long_factor = tuple(rope_scaling["long_factor"])
# Workaround for DeepSeek V3 Lite since its rope_scaling is null in config.json.
# Workaround for DeepSeek V3 Lite since its rope_scaling is null in
# config.json.
elif config.model_type == "deepseek_v3":
rope_params.scale_type = RotaryScalingType.yarn
# Other metdadata for RoPE.
Expand Down
24 changes: 22 additions & 2 deletions tensorrt_llm/_torch/auto_deploy/export/library/autocast_noop.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,16 +13,36 @@ class AutocastNoopPatch(BaseExportPatch):

This patch replaces torch.autocast with a null context manager
that can interfere with export.

It also patches ``torch.is_autocast_enabled`` so that transformers 5.x
helpers like ``maybe_autocast`` (used in RoPE embeddings) do not call the
real function with an unknown/fake device type during ``torch.export``
tracing, which would raise a ``RuntimeError``.
"""

def _apply_patch(self):
"""Apply the autocast no-op patch."""
# Store original function
# Store original functions
self.original_values["torch.autocast"] = torch.autocast
self.original_values["torch.is_autocast_enabled"] = torch.is_autocast_enabled

# Apply patch
# Apply patches
torch.autocast = lambda *args, **kwargs: nullcontext()

# torch.is_autocast_enabled(device_type) can fail during export when the
# device_type is unknown (e.g. fake/meta tensors). Return False so that
# callers like transformers' ``maybe_autocast`` skip the autocast block.
original_is_autocast = self.original_values["torch.is_autocast_enabled"]

def _safe_is_autocast_enabled(*args, **kwargs):
try:
return original_is_autocast(*args, **kwargs)
except RuntimeError:
return False

torch.is_autocast_enabled = _safe_is_autocast_enabled

def _revert_patch(self):
"""Revert the autocast no-op patch."""
torch.autocast = self.original_values["torch.autocast"]
torch.is_autocast_enabled = self.original_values["torch.is_autocast_enabled"]
Original file line number Diff line number Diff line change
@@ -0,0 +1,87 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Patch for transformers causal mask to be export-compatible with meta tensors.

Transformers 5.x's ``masking_utils.create_causal_mask`` calls
``find_packed_sequence_indices`` which invokes ``.all()`` on the attention mask
tensor. During ``torch.export`` tracing on meta tensors this raises because
``.item()`` / ``.all()`` are not supported on meta tensors.

This patch replaces ``find_packed_sequence_indices`` with a version that returns
early when the input tensors live on the meta device.
"""

import importlib.metadata

from packaging import version

from ..interface import BaseExportPatch, ExportPatchConfig, ExportPatchRegistry


def _transformers_version() -> str:
"""Get the version of transformers."""
return version.parse(importlib.metadata.version("transformers")).base_version


@ExportPatchRegistry.register("transformers_causal_mask")
class TransformersCausalMaskPatch(BaseExportPatch):
"""Patch ``find_packed_sequence_indices`` to handle meta tensors during export."""

@classmethod
def get_config_class(cls):
return ExportPatchConfig

def _apply_patch(self):
"""Apply the causal mask patch for meta-tensor compatibility."""
# Only needed for transformers >= 5.0.0 which introduced find_packed_sequence_indices
if version.parse(_transformers_version()) < version.parse("4.53.0"):
return

try:
from transformers import masking_utils

if not hasattr(masking_utils, "find_packed_sequence_indices"):
return

original_fn = masking_utils.find_packed_sequence_indices
self.original_values["find_packed_sequence_indices"] = original_fn

def _meta_safe_find_packed_sequence_indices(*args, **kwargs):
"""Wrapper that returns None for meta tensors, delegates otherwise."""
# The first positional arg is the attention_mask tensor
mask = args[0] if args else kwargs.get("attention_mask")
if mask is not None and hasattr(mask, "device") and mask.device.type == "meta":
return None
return original_fn(*args, **kwargs)

masking_utils.find_packed_sequence_indices = _meta_safe_find_packed_sequence_indices

except (ImportError, AttributeError):
pass

def _revert_patch(self):
"""Revert the causal mask patch."""
if version.parse(_transformers_version()) < version.parse("4.53.0"):
return

try:
from transformers import masking_utils

if "find_packed_sequence_indices" in self.original_values:
masking_utils.find_packed_sequence_indices = self.original_values[
"find_packed_sequence_indices"
]
except ImportError:
pass
Loading
Loading