Fixed new issue with conditions being on the wrong device

This commit is contained in:
Jaret Burkett
2025-08-05 13:40:09 -06:00
parent 8287164e9f
commit beae587bfa
+6
View File
@@ -163,8 +163,14 @@ def flex2_concat_cond(self: Flux, **kwargs):
def flex2_extra_conds(self, **kwargs):
out = self._flex2_orig_extra_conds(**kwargs)
noise = kwargs.get("noise", None)
device = kwargs["device"]
# needed now for some reason
for key in out.keys():
if hasattr(out[key], "cond"):
out[key].cond = out[key].cond.to(device)
flex2_concat_latent = kwargs.get("flex2_concat_latent", None)
flex2_concat_latent_no_control = kwargs.get(
"flex2_concat_latent_no_control", None)