Compare commits

...
Sign in to create a new pull request.

5 commits

Author SHA1 Message Date
Daniel Han
c45d5b6efa Fix compute_loss token_type_ids propagation and ModernBERT flex_attention
1. In grpo_trainer_compute_loss, extract token_type_ids and
   mm_token_type_ids from inputs dict and pass them to
   grpo_accumulated_loss. Without this, Gemma3 Vision GRPO fails with
   "token_type_ids is required as a model input when training" because
   the model forward never receives it.

2. In _generate_and_score_completions, extend mm_token_type_ids with
   zeros for completion tokens (parallel to existing token_type_ids
   handling). Also save mm_token_type_ids to the output dict. Without
   this, Qwen3VL GRPO fails with get_rope_index shape mismatch because
   mm_token_type_ids covers only the prompt, not prompt+completion.

3. Add "modernbert" to the flex_attention exclusion list in
   prefer_flex_attn_if_supported. ModernBERT with flex_attention hits a
   CUDA illegal memory access in create_block_mask's torch.compile path.
   Falling back to eager attention avoids the crash.

All changes are no-ops on transformers 4.x (token_type_ids/mm_token_type_ids
are None, and modernbert does not exist in the exclusion list check).
2026-04-01 09:13:10 +00:00
pre-commit-ci[bot]
d32ee870a7 [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
2026-04-01 08:32:58 +00:00
Daniel Han
2290407d99 Pass token_type_ids and mm_token_type_ids through GRPO VLM path
Transformers 5.x requires token_type_ids for some vision models during
training (e.g. Gemma3 Vision calls create_causal_mask_mapping which
raises ValueError if token_type_ids is None during training). Similarly,
Qwen3VL requires mm_token_type_ids for M-RoPE computation.

Extract both from kwargs in _get_per_token_logps_and_entropies, chunk
them alongside other vision tensors, and pass them to the model forward
call via _extra_vision_kwargs dict. This is a no-op when the tensors
are None (transformers 4.x or non-vision models).
2026-04-01 08:31:53 +00:00
pre-commit-ci[bot]
b506fbd86f [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
2026-04-01 07:28:57 +00:00
Daniel Han
c0cf5f257c Fix forward compatibility with transformers 5.x
Three issues fixed:

1. Skip exec-based config patching for transformers >= 5.0

Transformers 5.x config classes use @strict, @auto_docstring, and
interval() decorators/annotations that break exec(inspect.getsource(...)).
Those configs already use rope_parameters (the v5 replacement for
rope_scaling), so the patching is not needed. Gated with a version check
so transformers 4.x behavior is unchanged.

2. Slice position_ids to last token in fast_forward_inference

Transformers 5.x generate() accumulates position_ids as
[batch, full_seq_len] across decode steps instead of [batch, 1].
This causes a shape mismatch when indexing cos/sin for rotary
embeddings: cos[position_ids] produces [batch, full_seq_len, head_dim]
but Qn is [batch, n_heads, 1, head_dim]. Fixed by slicing
position_ids[:, -1:] when shape[-1] > 1. Applied to all model files
with fast_forward_inference: llama, qwen3, falcon_h1, gemma2, cohere,
granite. No-op on transformers 4.x since position_ids is already
[batch, 1]. Training path is unaffected.

3. Handle @strict config kwargs for sequence classification

Transformers 5.x @strict config decorator rejects unexpected kwargs
like num_labels, id2label, and max_position_embeddings passed to model
__init__(). Fixed by setting these on the config object directly and
passing config= to from_pretrained. Also added num_labels routing in
FastModel loader to select AutoModelForSequenceClassification.
2026-04-01 07:28:24 +00:00
10 changed files with 117 additions and 34 deletions

View file

@ -249,7 +249,7 @@ def prefer_flex_attn_if_supported(model_class, config):
# NemotronH: hybrid Mamba-2 + Transformer model that does not # NemotronH: hybrid Mamba-2 + Transformer model that does not
# support flex_attention (raises NotImplementedError from transformers). # support flex_attention (raises NotImplementedError from transformers).
model_type = getattr(config, "model_type", "") if config else "" model_type = getattr(config, "model_type", "") if config else ""
if model_type in ("gpt_oss", "mllama", "nemotron_h") or str( if model_type in ("gpt_oss", "mllama", "nemotron_h", "modernbert") or str(
model_type model_type
).startswith("gemma3n"): ).startswith("gemma3n"):
return None return None
@ -753,6 +753,11 @@ try:
except: except:
from transformers import PretrainedConfig from transformers import PretrainedConfig
# transformers 5.x uses class-level annotations + decorators (@strict, @auto_docstring, interval())
# in config classes, making exec(inspect.getsource(...)) infeasible. Skip config patching for 5.x
# since those configs already use rope_parameters (renamed from rope_scaling).
_skip_config_exec_patch = Version(transformers_version) >= Version("5.0.0")
model_architectures = [ model_architectures = [
"llama", "llama",
"mistral", "mistral",
@ -765,43 +770,44 @@ model_architectures = [
"falcon_h1", "falcon_h1",
] ]
for model_name in model_architectures: if not _skip_config_exec_patch:
config_filepath = f"transformers.models.{model_name}.configuration_{model_name}" for model_name in model_architectures:
model_filepath = f"transformers.models.{model_name}.modeling_{model_name}" config_filepath = f"transformers.models.{model_name}.configuration_{model_name}"
config_filename = f"{model_name.title().replace('_','')}Config" # qwen3 arch folder is qwen3_moe but config is Qwen3Config. Need to remove underscore(_) for now model_filepath = f"transformers.models.{model_name}.modeling_{model_name}"
try: config_filename = f"{model_name.title().replace('_','')}Config" # qwen3 arch folder is qwen3_moe but config is Qwen3Config. Need to remove underscore(_) for now
exec(f"from {config_filepath} import {config_filename}", globals())
except:
continue
try:
config = inspect.getsource(eval(config_filename))
except:
continue
if "RopeParameters" in config:
try: try:
exec(f"from {config_filepath} import RopeParameters", globals()) exec(f"from {config_filepath} import {config_filename}", globals())
except: except:
continue continue
if "rope_scaling" in config: try:
continue config = inspect.getsource(eval(config_filename))
config = re.sub( except:
r"(\*\*kwargs)[\s]{0,}\,[\s]{0,}\)[\s]{0,}\:", continue
r"rope_scaling=None," if "RopeParameters" in config:
r"\n **kwargs):\n" try:
r"\n self.rope_scaling = rope_scaling\n", exec(f"from {config_filepath} import RopeParameters", globals())
config, except:
) continue
# Just for Mistral Nemo if "rope_scaling" in config:
if model_name == "mistral": continue
if Version(transformers_version) <= Version("4.42.4"): config = re.sub(
config = patch_mistral_nemo_config(config) r"(\*\*kwargs)[\s]{0,}\,[\s]{0,}\)[\s]{0,}\:",
r"rope_scaling=None,"
r"\n **kwargs):\n"
r"\n self.rope_scaling = rope_scaling\n",
config,
)
exec(config, globals()) # Just for Mistral Nemo
exec(f"import {config_filepath}", globals()) if model_name == "mistral":
exec(f"{config_filepath}.{config_filename} = {config_filename}", globals()) if Version(transformers_version) <= Version("4.42.4"):
config = patch_mistral_nemo_config(config)
exec(config, globals())
exec(f"import {config_filepath}", globals())
exec(f"{config_filepath}.{config_filename} = {config_filename}", globals())
# ============================================= # =============================================
# ============================================= # =============================================

View file

@ -357,6 +357,9 @@ def CohereAttention_fast_forward_inference(
# cos, sin = self.rotary_emb(Vn, seq_len = kv_seq_len) # cos, sin = self.rotary_emb(Vn, seq_len = kv_seq_len)
# Qn, Kn = inplace_rope_embedding(Qn, Kn, cos, sin, position_ids) # Qn, Kn = inplace_rope_embedding(Qn, Kn, cos, sin, position_ids)
cos, sin = self.rotary_emb.get_cached(kv_seq_len, Qn.device.index) cos, sin = self.rotary_emb.get_cached(kv_seq_len, Qn.device.index)
# Transformers 5.x: position_ids may be [batch, full_seq_len]; slice to last
if position_ids.dim() >= 2 and position_ids.shape[-1] > 1:
position_ids = position_ids[:, -1:]
cos = cos[position_ids].unsqueeze(1) cos = cos[position_ids].unsqueeze(1)
sin = sin[position_ids].unsqueeze(1) sin = sin[position_ids].unsqueeze(1)
h = self.half_head_dim h = self.half_head_dim

View file

@ -313,6 +313,9 @@ def FalconH1Attention_fast_forward_inference(
# or else error # or else error
self.rotary_emb.extend_rope_embedding(Vn, seq_len + 2) self.rotary_emb.extend_rope_embedding(Vn, seq_len + 2)
cos, sin = self.rotary_emb.get_cached(kv_seq_len, Qn.device.index) cos, sin = self.rotary_emb.get_cached(kv_seq_len, Qn.device.index)
# Transformers 5.x: position_ids may be [batch, full_seq_len]; slice to last
if position_ids.dim() >= 2 and position_ids.shape[-1] > 1:
position_ids = position_ids[:, -1:]
cos = cos[position_ids].unsqueeze(1) cos = cos[position_ids].unsqueeze(1)
sin = sin[position_ids].unsqueeze(1) sin = sin[position_ids].unsqueeze(1)
h = self.half_head_dim h = self.half_head_dim

View file

@ -394,6 +394,9 @@ def Gemma2Attention_fast_forward_inference(
# cos, sin = self.rotary_emb(Vn, seq_len = kv_seq_len) # cos, sin = self.rotary_emb(Vn, seq_len = kv_seq_len)
# Qn, Kn = inplace_rope_embedding(Qn, Kn, cos, sin, position_ids) # Qn, Kn = inplace_rope_embedding(Qn, Kn, cos, sin, position_ids)
cos, sin = self.rotary_emb.get_cached(kv_seq_len, Qn.device.index) cos, sin = self.rotary_emb.get_cached(kv_seq_len, Qn.device.index)
# Transformers 5.x: position_ids may be [batch, full_seq_len]; slice to last
if position_ids.dim() >= 2 and position_ids.shape[-1] > 1:
position_ids = position_ids[:, -1:]
cos = cos[position_ids].unsqueeze(1) cos = cos[position_ids].unsqueeze(1)
sin = sin[position_ids].unsqueeze(1) sin = sin[position_ids].unsqueeze(1)
h = self.half_head_dim h = self.half_head_dim

View file

@ -355,6 +355,9 @@ def GraniteAttention_fast_forward_inference(
# cos, sin = self.rotary_emb(Vn, seq_len = kv_seq_len) # cos, sin = self.rotary_emb(Vn, seq_len = kv_seq_len)
# Qn, Kn = inplace_rope_embedding(Qn, Kn, cos, sin, position_ids) # Qn, Kn = inplace_rope_embedding(Qn, Kn, cos, sin, position_ids)
cos, sin = position_embeddings cos, sin = position_embeddings
# Transformers 5.x: position_ids may be [batch, full_seq_len]; slice to last
if position_ids.dim() >= 2 and position_ids.shape[-1] > 1:
position_ids = position_ids[:, -1:]
cos, sin = cos[position_ids], sin[position_ids] cos, sin = cos[position_ids], sin[position_ids]
h = self.half_head_dim h = self.half_head_dim

View file

@ -496,6 +496,10 @@ def LlamaAttention_fast_forward_inference(
# ensure correct shape # ensure correct shape
if position_ids.dim() == 1: if position_ids.dim() == 1:
position_ids = position_ids[:, None] position_ids = position_ids[:, None]
# Transformers 5.x generate() accumulates position_ids as [batch, full_seq_len]
# across decode steps. In single-token inference we only need the last position.
if position_ids.shape[-1] > 1:
position_ids = position_ids[:, -1:]
position_ids = position_ids.to(Qn.device) position_ids = position_ids.to(Qn.device)
if rotary_seq_len is None: if rotary_seq_len is None:
@ -2414,14 +2418,19 @@ class FastLlamaModel:
raise_handler = RaiseUninitialized() raise_handler = RaiseUninitialized()
if num_labels is not None: if num_labels is not None:
# Transformers 5.x @strict config classes reject unexpected kwargs
# like num_labels and max_position_embeddings. Set on the config
# object directly and pass config= instead.
model_config.num_labels = num_labels
if max_position_embeddings is not None:
model_config.max_position_embeddings = max_position_embeddings
model = AutoModelForSequenceClassification.from_pretrained( model = AutoModelForSequenceClassification.from_pretrained(
model_name, model_name,
config = model_config,
device_map = device_map, device_map = device_map,
# torch_dtype = dtype, # transformers changed torch_dtype to dtype # torch_dtype = dtype, # transformers changed torch_dtype to dtype
num_labels = num_labels,
# quantization_config = bnb_config, # quantization_config = bnb_config,
token = token, token = token,
max_position_embeddings = max_position_embeddings,
trust_remote_code = trust_remote_code, trust_remote_code = trust_remote_code,
attn_implementation = preferred_attn_impl, attn_implementation = preferred_attn_impl,
**kwargs, **kwargs,

View file

@ -1407,8 +1407,14 @@ class FastModel(FastBaseModel):
architectures = [] architectures = []
is_vlm = any(x.endswith("ForConditionalGeneration") for x in architectures) is_vlm = any(x.endswith("ForConditionalGeneration") for x in architectures)
is_vlm = is_vlm or hasattr(model_config, "vision_config") is_vlm = is_vlm or hasattr(model_config, "vision_config")
# If num_labels is set, use AutoModelForSequenceClassification
_num_labels = kwargs.get("num_labels", None)
if auto_model is None: if auto_model is None:
if is_vlm: if _num_labels is not None:
from transformers import AutoModelForSequenceClassification
auto_model = AutoModelForSequenceClassification
elif is_vlm:
# Check if the model's auto_map supports the VLM auto class. # Check if the model's auto_map supports the VLM auto class.
# Some VL models (e.g. Nemotron-VL) only register AutoModelForCausalLM # Some VL models (e.g. Nemotron-VL) only register AutoModelForCausalLM
# in their auto_map, not AutoModelForImageTextToText/AutoModelForVision2Seq. # in their auto_map, not AutoModelForImageTextToText/AutoModelForVision2Seq.

View file

@ -302,6 +302,9 @@ def Qwen3Attention_fast_forward_inference(
# or else error # or else error
self.rotary_emb.extend_rope_embedding(Vn, seq_len + 2) self.rotary_emb.extend_rope_embedding(Vn, seq_len + 2)
cos, sin = self.rotary_emb.get_cached(kv_seq_len, Qn.device.index) cos, sin = self.rotary_emb.get_cached(kv_seq_len, Qn.device.index)
# Transformers 5.x: position_ids may be [batch, full_seq_len]; slice to last
if position_ids.dim() >= 2 and position_ids.shape[-1] > 1:
position_ids = position_ids[:, -1:]
cos = cos[position_ids].unsqueeze(1) cos = cos[position_ids].unsqueeze(1)
sin = sin[position_ids].unsqueeze(1) sin = sin[position_ids].unsqueeze(1)
h = self.half_head_dim h = self.half_head_dim

View file

@ -542,6 +542,17 @@ def grpo_trainer__generate_and_score_completions(function_name, function):
function = patched function = patched
# Transformers 5.x: Extend mm_token_type_ids for completion tokens (Qwen3VL M-RoPE)
# TRL handles token_type_ids but not mm_token_type_ids
_tt_search = 'if "token_type_ids" in forward_kwargs:\n token_type_ids = forward_kwargs["token_type_ids"]\n forward_kwargs["token_type_ids"] = torch.cat(\n [token_type_ids, token_type_ids.new_zeros(completion_ids.shape)], dim=1\n )'
_tt_replace = _tt_search + '\n if "mm_token_type_ids" in forward_kwargs:\n mm_tti = forward_kwargs["mm_token_type_ids"]\n forward_kwargs["mm_token_type_ids"] = torch.cat(\n [mm_tti, mm_tti.new_zeros(completion_ids.shape)], dim=1\n )'
function = function.replace(_tt_search, _tt_replace)
# Save mm_token_type_ids to output dict alongside token_type_ids
_save_search = 'if "token_type_ids" in forward_kwargs:\n output["token_type_ids"] = forward_kwargs["token_type_ids"]'
_save_replace = _save_search + '\n if "mm_token_type_ids" in forward_kwargs:\n output["mm_token_type_ids"] = forward_kwargs["mm_token_type_ids"]'
function = function.replace(_save_search, _save_replace)
return function return function
@ -714,6 +725,9 @@ def grpo_trainer__get_per_token_logps_and_entropies(function_name, function):
kwargs.get("pixel_attention_mask", None), kwargs.get("pixel_attention_mask", None),
kwargs.get("image_sizes", None), kwargs.get("image_sizes", None),
) )
# Transformers 5.x needs token_type_ids/mm_token_type_ids for some vision models
token_type_ids = kwargs.get("token_type_ids", None)
mm_token_type_ids = kwargs.get("mm_token_type_ids", None)
unwrapped_model = self.accelerator.unwrap_model( unwrapped_model = self.accelerator.unwrap_model(
model, keep_fp32_wrapper = False model, keep_fp32_wrapper = False
@ -831,6 +845,10 @@ def grpo_trainer__get_per_token_logps_and_entropies(function_name, function):
if logit_scale_divide is None: if logit_scale_divide is None:
logit_scale_divide = 0 logit_scale_divide = 0
# Transformers 5.x needs token_type_ids/mm_token_type_ids for some vision models
token_type_ids_chunks = chunk_optional(token_type_ids, B)
mm_token_type_ids_chunks = chunk_optional(mm_token_type_ids, B)
zipped_inputs = zip( zipped_inputs = zip(
input_ids_chunks, input_ids_chunks,
attention_mask_chunks, attention_mask_chunks,
@ -838,6 +856,8 @@ def grpo_trainer__get_per_token_logps_and_entropies(function_name, function):
image_grid_thw_chunks, image_grid_thw_chunks,
pixel_attention_mask_chunks, pixel_attention_mask_chunks,
image_sizes_chunks, image_sizes_chunks,
token_type_ids_chunks,
mm_token_type_ids_chunks,
) )
os.environ["UNSLOTH_RETURN_HIDDEN_STATES"] = "1" os.environ["UNSLOTH_RETURN_HIDDEN_STATES"] = "1"
@ -849,7 +869,16 @@ def grpo_trainer__get_per_token_logps_and_entropies(function_name, function):
image_grid_thw_chunk, image_grid_thw_chunk,
pixel_attention_mask_chunk, pixel_attention_mask_chunk,
image_sizes_chunk, image_sizes_chunk,
token_type_ids_chunk,
mm_token_type_ids_chunk,
) in zipped_inputs: ) in zipped_inputs:
_extra_vision_kwargs = {}
if token_type_ids_chunk is not None:
_extra_vision_kwargs["token_type_ids"] = token_type_ids_chunk
if mm_token_type_ids_chunk is not None:
_extra_vision_kwargs["mm_token_type_ids"] = (
mm_token_type_ids_chunk
)
with torch.amp.autocast( with torch.amp.autocast(
device_type = "cuda", dtype = self._autocast_dtype device_type = "cuda", dtype = self._autocast_dtype
): ):
@ -861,6 +890,7 @@ def grpo_trainer__get_per_token_logps_and_entropies(function_name, function):
image_grid_thw = image_grid_thw_chunk, image_grid_thw = image_grid_thw_chunk,
pixel_attention_mask = pixel_attention_mask_chunk, pixel_attention_mask = pixel_attention_mask_chunk,
image_sizes = image_sizes_chunk, image_sizes = image_sizes_chunk,
**_extra_vision_kwargs,
).logits ).logits
completion_input_ids_chunk = input_ids_chunk[ completion_input_ids_chunk = input_ids_chunk[
@ -893,6 +923,7 @@ def grpo_trainer__get_per_token_logps_and_entropies(function_name, function):
pixel_attention_mask = pixel_attention_mask_chunk, pixel_attention_mask = pixel_attention_mask_chunk,
image_sizes = image_sizes_chunk, image_sizes = image_sizes_chunk,
logits_to_keep = logits_to_keep + 1, logits_to_keep = logits_to_keep + 1,
**_extra_vision_kwargs,
).logits ).logits
logits_chunk = logits_chunk[:, :-1, :] logits_chunk = logits_chunk[:, :-1, :]
@ -993,6 +1024,9 @@ def grpo_trainer_compute_loss(function_name, function):
inputs.get("pixel_attention_mask", None), inputs.get("pixel_attention_mask", None),
inputs.get("image_sizes", None), inputs.get("image_sizes", None),
) )
# Transformers 5.x needs token_type_ids/mm_token_type_ids for some vision models
token_type_ids = inputs.get("token_type_ids", None)
mm_token_type_ids = inputs.get("mm_token_type_ids", None)
num_items_in_batch = inputs.get("num_items_in_batch", None) num_items_in_batch = inputs.get("num_items_in_batch", None)
sampling_per_token_logps = inputs.get("sampling_per_token_logps", None) sampling_per_token_logps = inputs.get("sampling_per_token_logps", None)
current_gradient_accumulation_steps = self.current_gradient_accumulation_steps current_gradient_accumulation_steps = self.current_gradient_accumulation_steps
@ -1136,6 +1170,8 @@ def grpo_trainer_compute_loss(function_name, function):
current_gradient_accumulation_steps = current_gradient_accumulation_steps, current_gradient_accumulation_steps = current_gradient_accumulation_steps,
num_processes = num_processes, num_processes = num_processes,
sampling_per_token_logps = sampling_per_token_logps, sampling_per_token_logps = sampling_per_token_logps,
token_type_ids = token_type_ids,
mm_token_type_ids = mm_token_type_ids,
) )
else: else:
# to ensure backwards compatibility with trl 0.15.2 and maybe even 0.17 # to ensure backwards compatibility with trl 0.15.2 and maybe even 0.17
@ -1154,6 +1190,8 @@ def grpo_trainer_compute_loss(function_name, function):
logit_scale_multiply = logit_scale_multiply, logit_scale_multiply = logit_scale_multiply,
logit_scale_divide = logit_scale_divide, logit_scale_divide = logit_scale_divide,
attention_mask = attention_mask, attention_mask = attention_mask,
token_type_ids = token_type_ids,
mm_token_type_ids = mm_token_type_ids,
) )
) )
if "train" in self._metrics: if "train" in self._metrics:

View file

@ -788,6 +788,15 @@ class FastBaseModel:
if not fast_inference: if not fast_inference:
# Prevent load_in_fp8 from being forwarded into HF internal model loading # Prevent load_in_fp8 from being forwarded into HF internal model loading
load_in_fp8 = kwargs.pop("load_in_fp8", None) load_in_fp8 = kwargs.pop("load_in_fp8", None)
# Transformers 5.x @strict config classes reject unexpected kwargs.
# Move config-level attributes onto the config object directly.
_num_labels = kwargs.pop("num_labels", None)
if _num_labels is not None:
model_config.num_labels = _num_labels
for _cfg_key in ("id2label", "label2id", "max_position_embeddings"):
_cfg_val = kwargs.pop(_cfg_key, None)
if _cfg_val is not None:
setattr(model_config, _cfg_key, _cfg_val)
model = auto_model.from_pretrained( model = auto_model.from_pretrained(
model_name, model_name,
config = model_config, config = model_config,