From 8c89aafa56ad970ae9efe7c097cda80892eacc3e Mon Sep 17 00:00:00 2001 From: rockdu Date: Sat, 8 Aug 2026 01:31:48 -0700 Subject: [PATCH] fix(sft): drop the cond cast deleted with loss_hub.utils --- miles/backends/fsdp_utils/loss_hub/flow_grpo.py | 1 - miles/backends/fsdp_utils/loss_hub/nft.py | 1 - miles/backends/fsdp_utils/loss_hub/sft.py | 6 +----- 3 files changed, 1 insertion(+), 7 deletions(-) diff --git a/miles/backends/fsdp_utils/loss_hub/flow_grpo.py b/miles/backends/fsdp_utils/loss_hub/flow_grpo.py index 876d69dc..39a69d18 100644 --- a/miles/backends/fsdp_utils/loss_hub/flow_grpo.py +++ b/miles/backends/fsdp_utils/loss_hub/flow_grpo.py @@ -73,7 +73,6 @@ def prepare_flow_grpo_batch( if use_cfg else None ) - # Cond dtypes are set at the model boundary by the family input_dtype_policy (see actor). cfg_batching = use_cfg and bool(args.fsdp_cfg_batching) joint_cond = pos_cond = neg_cond = None if cfg_batching: diff --git a/miles/backends/fsdp_utils/loss_hub/nft.py b/miles/backends/fsdp_utils/loss_hub/nft.py index 97932078..a4f46f4c 100644 --- a/miles/backends/fsdp_utils/loss_hub/nft.py +++ b/miles/backends/fsdp_utils/loss_hub/nft.py @@ -37,7 +37,6 @@ def prepare_nft_batch( component_name, model = next(iter(ctx.models.items())) pos_list = [config.prepare_cond_kwargs(batch[i]["denoising_env"].pos_cond_kwargs, device) for i in range(bsz)] - # Cond dtypes are set at the model boundary by the family input_dtype_policy (see actor). pos_cond = config.collate_cond_for_sample_batch(pos_list, device, pad_to_len=pad_to_len) num_train_timesteps = ctx.scheduler.config.num_train_timesteps diff --git a/miles/backends/fsdp_utils/loss_hub/sft.py b/miles/backends/fsdp_utils/loss_hub/sft.py index fd850941..331998b4 100644 --- a/miles/backends/fsdp_utils/loss_hub/sft.py +++ b/miles/backends/fsdp_utils/loss_hub/sft.py @@ -8,7 +8,6 @@ import torch.nn as nn from miles.backends.fsdp_utils.loss_hub.types import DiffusionLossContext, PreparedBatch -from miles.backends.fsdp_utils.loss_hub.utils import cast_cond_to_dtype from miles.utils.metric_buffer import MetricBuffer @@ -90,10 +89,7 @@ def prepare_sft_batch( timesteps_for_model = timesteps cond_list = [{key: value.to(device) for key, value in pair["cond_kwargs"].items()} for pair in batch] - pos_cond = cast_cond_to_dtype( - config.collate_cond_for_sample_batch(cond_list, device, pad_to_len=pad_to_len), - ctx.forward_dtype, - ) + pos_cond = config.collate_cond_for_sample_batch(cond_list, device, pad_to_len=pad_to_len) return PreparedBatch( latents=latents,