Fixed new issue with conditions being on the wrong device
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user