diff --git a/comfy/ldm/mage_flow/model.py b/comfy/ldm/mage_flow/model.py index 92a5faa52..ac29bb610 100644 --- a/comfy/ldm/mage_flow/model.py +++ b/comfy/ldm/mage_flow/model.py @@ -22,7 +22,6 @@ class MageTimestepProjEmbeddings(nn.Module): ) def forward(self, timestep, hidden_states): - timestep = timestep.to(hidden_states.dtype) half_dim = 128 exponent = -math.log(10000) * torch.arange(half_dim, dtype=torch.float32, device=timestep.device) / half_dim emb = torch.exp(exponent).to(timestep.dtype) diff --git a/comfy/model_base.py b/comfy/model_base.py index 50c73a431..ee6dc57a2 100644 --- a/comfy/model_base.py +++ b/comfy/model_base.py @@ -2279,6 +2279,10 @@ class MageFlow(QwenImage): def __init__(self, model_config, model_type=ModelType.FLOW, device=None): super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.mage_flow.model.MageFlowTransformer2DModel) + def process_timestep(self, timestep, **kwargs): + # Mage runs in bf16 and rounds its timestep frequency table to the timestep dtype, keep that on fp32 devices. + return timestep.to(torch.bfloat16) + def extra_conds_shapes(self, **kwargs): out = {} ref_latents = kwargs.get("reference_latents", None)