NguyenDinhHieu commited on
Commit
fb2fb6b
·
verified ·
1 Parent(s): c9aa3db

Update cldm/cldm.py

Browse files
Files changed (1) hide show
  1. 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) # 接收t
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
- # 设置条件;根据返回的t;设置输入的条件图;高层的条件图为全黑的
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] # 如果定义了bs,则截取前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
- # 重写的get_input也需要返回t
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):