From 8d356ebb398feb49316280a70752e049ce2cf780 Mon Sep 17 00:00:00 2001 From: dgouju Date: Wed, 29 Jul 2026 18:39:14 +0200 Subject: [PATCH] Add images in SFT loss_func for multimodal post training MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit MaxText's post-training wrapper function (use_maxtext_loss_function in train_sft.py) has a hardcoded function signature accepting only text keyword arguments (inputs, targets, etc.), but Tunix's batch iterator unpacks all batch dictionary keys—including images and image_masks—into the loss function during multimodal training. --- src/maxtext/trainers/post_train/sft/train_sft.py | 7 +++++++ 1 file changed, 7 insertions(+) diff --git a/src/maxtext/trainers/post_train/sft/train_sft.py b/src/maxtext/trainers/post_train/sft/train_sft.py index 9002caa16c..e1cf97c62d 100644 --- a/src/maxtext/trainers/post_train/sft/train_sft.py +++ b/src/maxtext/trainers/post_train/sft/train_sft.py @@ -210,6 +210,9 @@ def loss_func( targets, targets_position, targets_segmentation, + images=None, + image_masks=None, + **kwargs, ): data = { "inputs": inputs, @@ -219,6 +222,10 @@ def loss_func( "targets_position": targets_position, "targets_segmentation": targets_segmentation, } + if images is not None: + data["images"] = images + if image_masks is not None: + data["image_masks"] = image_masks return loss_fn(model, mt_config, data, dropout_rng=None, params=None, is_train=True) trainer = trainer.with_loss_fn(loss_func, has_aux=True)