Spaces:
Paused
Paused
Update cldm/cldm.py
Browse files- cldm/cldm.py +4 -6
cldm/cldm.py
CHANGED
|
@@ -320,14 +320,14 @@ class ControlLDM(LatentDiffusion):
|
|
| 320 |
@torch.no_grad()
|
| 321 |
def get_input(self, batch, k, bs=None, *args, **kwargs):
|
| 322 |
# x, c = super().get_input(batch, self.first_stage_key, *args, **kwargs)
|
| 323 |
-
x, c, t = super().get_input(batch, self.first_stage_key, *args, **kwargs)
|
| 324 |
control = batch[self.control_key]
|
| 325 |
if bs is not None:
|
| 326 |
control = control[:bs]
|
| 327 |
control = control.to(self.device)
|
| 328 |
control = einops.rearrange(control, 'b h w c -> b c h w')
|
| 329 |
control = control.to(memory_format=torch.contiguous_format).float()
|
| 330 |
-
|
| 331 |
# control_processed=[]
|
| 332 |
# for i in range(len(t)):
|
| 333 |
# control = batch[self.control_key][i].unsqueeze(0)
|
|
@@ -336,7 +336,7 @@ class ControlLDM(LatentDiffusion):
|
|
| 336 |
# black_image = torch.zeros(control_shape, dtype=torch.float32, device=self.device)
|
| 337 |
# control = black_image
|
| 338 |
# if bs is not None:
|
| 339 |
-
# control = control[:bs]
|
| 340 |
# control = control.to(self.device)
|
| 341 |
# control = einops.rearrange(control, 'b h w c -> b c h w')
|
| 342 |
# control = control.to(memory_format=torch.contiguous_format).float()
|
|
@@ -349,10 +349,8 @@ class ControlLDM(LatentDiffusion):
|
|
| 349 |
# control = control.to(memory_format=torch.contiguous_format).float()
|
| 350 |
# control_processed.append(control)
|
| 351 |
|
| 352 |
-
# # 将列表中的张量拼接在一起,形成一个新的张量,沿着新的第一个维度
|
| 353 |
# control = torch.cat(control_processed, dim=0)
|
| 354 |
-
|
| 355 |
-
# 进入ddpm中的shared_setp()
|
| 356 |
return x, dict(c_crossattn=[c], c_concat=[control]), t
|
| 357 |
|
| 358 |
def apply_model(self, x_noisy, t, cond, *args, **kwargs):
|
|
|
|
| 320 |
@torch.no_grad()
|
| 321 |
def get_input(self, batch, k, bs=None, *args, **kwargs):
|
| 322 |
# x, c = super().get_input(batch, self.first_stage_key, *args, **kwargs)
|
| 323 |
+
x, c, t = super().get_input(batch, self.first_stage_key, *args, **kwargs)
|
| 324 |
control = batch[self.control_key]
|
| 325 |
if bs is not None:
|
| 326 |
control = control[:bs]
|
| 327 |
control = control.to(self.device)
|
| 328 |
control = einops.rearrange(control, 'b h w c -> b c h w')
|
| 329 |
control = control.to(memory_format=torch.contiguous_format).float()
|
| 330 |
+
|
| 331 |
# control_processed=[]
|
| 332 |
# for i in range(len(t)):
|
| 333 |
# control = batch[self.control_key][i].unsqueeze(0)
|
|
|
|
| 336 |
# black_image = torch.zeros(control_shape, dtype=torch.float32, device=self.device)
|
| 337 |
# control = black_image
|
| 338 |
# if bs is not None:
|
| 339 |
+
# control = control[:bs]
|
| 340 |
# control = control.to(self.device)
|
| 341 |
# control = einops.rearrange(control, 'b h w c -> b c h w')
|
| 342 |
# control = control.to(memory_format=torch.contiguous_format).float()
|
|
|
|
| 349 |
# control = control.to(memory_format=torch.contiguous_format).float()
|
| 350 |
# control_processed.append(control)
|
| 351 |
|
|
|
|
| 352 |
# control = torch.cat(control_processed, dim=0)
|
| 353 |
+
|
|
|
|
| 354 |
return x, dict(c_crossattn=[c], c_concat=[control]), t
|
| 355 |
|
| 356 |
def apply_model(self, x_noisy, t, cond, *args, **kwargs):
|