Compare commits
5 commits
main
...
fix/transf
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c45d5b6efa | ||
|
|
d32ee870a7 | ||
|
|
2290407d99 | ||
|
|
b506fbd86f | ||
|
|
c0cf5f257c |
10 changed files with 117 additions and 34 deletions
|
|
@ -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())
|
||||||
# =============================================
|
# =============================================
|
||||||
|
|
||||||
# =============================================
|
# =============================================
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
|
|
|
||||||
|
|
@ -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.
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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:
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue