From d2332622d1e6f2ca9a273ea4f7e706b851d5aad0 Mon Sep 17 00:00:00 2001 From: Roland Tannous Date: Tue, 17 Feb 2026 23:49:13 +0000 Subject: [PATCH] fix: defensively rename VLM chat column to match model's forward() signature --- studio/backend/core/training/trainer.py | 2 + .../backend/utils/datasets/dataset_utils.py | 23 +++++++++++- .../utils/datasets/format_conversion.py | 37 +++++++++++++++++++ 3 files changed, 61 insertions(+), 1 deletion(-) diff --git a/studio/backend/core/training/trainer.py b/studio/backend/core/training/trainer.py index 9468e282ef..a3abd4bd17 100644 --- a/studio/backend/core/training/trainer.py +++ b/studio/backend/core/training/trainer.py @@ -425,6 +425,7 @@ class UnslothTrainer: dataset, model_name=self.model_name, tokenizer=self.tokenizer, + model=self.model, is_vlm=self.is_vlm, format_type=format_type, dataset_name=dataset_source, @@ -447,6 +448,7 @@ class UnslothTrainer: eval_dataset, model_name=self.model_name, tokenizer=self.tokenizer, + model=self.model, is_vlm=self.is_vlm, format_type=format_type, dataset_name=dataset_source, diff --git a/studio/backend/utils/datasets/dataset_utils.py b/studio/backend/utils/datasets/dataset_utils.py index a75f78d37c..449d96d392 100644 --- a/studio/backend/utils/datasets/dataset_utils.py +++ b/studio/backend/utils/datasets/dataset_utils.py @@ -28,6 +28,8 @@ from .format_conversion import ( convert_alpaca_to_chatml, convert_to_vlm_format, convert_llava_to_vlm_format, + get_expected_chat_column, + rename_chat_column_in_list, ) from .chat_templates import ( apply_chat_template_to_dataset, @@ -547,6 +549,7 @@ def format_and_template_dataset( dataset, model_name, tokenizer, + model=None, is_vlm = False, format_type="auto", # VLM-specific parameters @@ -735,12 +738,30 @@ def format_and_template_dataset( dataset = [sample for sample in dataset] warnings.append("Dataset already in standard VLM messages format") + # Defensive: rename chat column if model expects a different name + expected_col = get_expected_chat_column(model) if model is not None else None + # VLM data is a list of dicts — check what key the first item uses + current_col = "messages" # default from our converters + if isinstance(dataset, list) and len(dataset) > 0: + sample_keys = dataset[0].keys() + if "conversations" in sample_keys: + current_col = "conversations" + elif "messages" in sample_keys: + current_col = "messages" + + if expected_col and expected_col != current_col: + warnings.append( + f"Model expects '{expected_col}' but dataset has '{current_col}' — renaming." + ) + dataset = rename_chat_column_in_list(dataset, current_col, expected_col) + current_col = expected_col + # Return as list return { "dataset": dataset, "detected_format": vlm_structure["format"], "final_format": "vlm_messages", - "chat_column": "messages", + "chat_column": current_col, "is_vlm": True, "is_multimodal": multimodal_info["is_multimodal"], "multimodal_info": multimodal_info, diff --git a/studio/backend/utils/datasets/format_conversion.py b/studio/backend/utils/datasets/format_conversion.py index 9367741e8e..841fde0fd6 100644 --- a/studio/backend/utils/datasets/format_conversion.py +++ b/studio/backend/utils/datasets/format_conversion.py @@ -8,6 +8,43 @@ This module contains functions for converting between dataset formats from datasets import IterableDataset +def get_expected_chat_column(model): + """ + Inspect the model's forward() signature to determine if it expects + 'messages' or 'conversations' as a column name. + + Returns: + str or None: 'messages', 'conversations', or None if neither found. + """ + import inspect + try: + sig = inspect.signature(model.forward) + params = list(sig.parameters.keys()) + + if "messages" in params: + return "messages" + elif "conversations" in params: + return "conversations" + except (ValueError, TypeError): + pass + return None + + +def rename_chat_column_in_list(data, from_col, to_col): + """ + Rename a chat column key in a list of dicts (for VLM data). + """ + if from_col == to_col: + return data + renamed = [] + for item in data: + new_item = {} + for k, v in item.items(): + new_item[to_col if k == from_col else k] = v + renamed.append(new_item) + return renamed + + def standardize_chat_format( dataset, tokenizer=None,