Fix MageFlow on cards that don't support bf16 (#15081)
This commit is contained in:
parent
f966a2b38c
commit
806e092ed4
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Reference in New Issue