Skip to content

Commit 12b43d0

Browse files
committed
locate the text model by config instead of by module name
1 parent d3618f2 commit 12b43d0

1 file changed

Lines changed: 11 additions & 3 deletions

File tree

trl/experimental/async_grpo/async_grpo_trainer.py

Lines changed: 11 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -33,7 +33,7 @@
3333
from datasets import Dataset, IterableDataset
3434
from torch.distributed._tensor import DTensor
3535
from torch.utils.data import DataLoader
36-
from transformers import AutoTokenizer, PreTrainedTokenizerBase, TrainerCallback
36+
from transformers import AutoTokenizer, PreTrainedModel, PreTrainedTokenizerBase, TrainerCallback
3737
from transformers.data.data_collator import DataCollatorMixin
3838
from transformers.trainer_utils import PREFIX_CHECKPOINT_DIR
3939

@@ -826,8 +826,16 @@ def __init__(
826826

827827
self._is_vlm = text_config is not model.config
828828
if self._is_vlm:
829-
model.model.requires_grad_(False)
830-
model.model.language_model.requires_grad_(True)
829+
# Train the text model only. It is located through the text config, since module names differ across
830+
# architectures (`model.language_model` for Qwen-VL and Gemma 3, `model.text_model` for SmolVLM).
831+
text_model = next(
832+
module
833+
for module in model.modules()
834+
if isinstance(module, PreTrainedModel) and module is not model and module.config is text_config
835+
)
836+
model.requires_grad_(False)
837+
text_model.requires_grad_(True)
838+
model.get_output_embeddings().requires_grad_(True)
831839

832840
patch_chunked_lm_head(
833841
model, chunk_size=8192, temperature=self.temperature, output_router_logits=self.aux_loss_enabled

0 commit comments

Comments
 (0)