diff --git a/gemma/diffusion/hackable_diffusion_adapter/hd/sft_model.py b/gemma/diffusion/hackable_diffusion_adapter/hd/sft_model.py index eb829227..d98ba017 100644 --- a/gemma/diffusion/hackable_diffusion_adapter/hd/sft_model.py +++ b/gemma/diffusion/hackable_diffusion_adapter/hd/sft_model.py @@ -404,7 +404,7 @@ def __call__( jax.random.uniform(self.make_rng('sampling'), shape=(batch_size,)) < self.self_cond_prob ) - # Reshape to broadcast with x0_hat_logits (Batch, ..., Channels) + # Reshape to broadcast with sc_logits (Batch, ..., Channels) do_self_cond = do_self_cond.reshape( (batch_size,) + (1,) * (sc_logits.ndim - 1) )