MCPcopy Create free account
hub / github.com/OpenDriveLab/ReSim / forward

Method forward

sat/sgm/modules/diffusionmodules/model.py:381–426  ·  view source on GitHub ↗
(self, x, t=None, context=None)

Source from the content-addressed store, hash-verified

379 self.conv_out = torch.nn.Conv2d(block_in, out_ch, kernel_size=3, stride=1, padding=1)
380
381 def forward(self, x, t=None, context=None):
382 # assert x.shape[2] == x.shape[3] == self.resolution
383 if context is not None:
384 # assume aligned context, cat along channel axis
385 x = torch.cat((x, context), dim=1)
386 if self.use_timestep:
387 # timestep embedding
388 assert t is not None
389 temb = get_timestep_embedding(t, self.ch)
390 temb = self.temb.dense[0](temb)
391 temb = nonlinearity(temb)
392 temb = self.temb.dense[1](temb)
393 else:
394 temb = None
395
396 # downsampling
397 hs = [self.conv_in(x)]
398 for i_level in range(self.num_resolutions):
399 for i_block in range(self.num_res_blocks):
400 h = self.down[i_level].block[i_block](hs[-1], temb)
401 if len(self.down[i_level].attn) > 0:
402 h = self.down[i_level].attn[i_block](h)
403 hs.append(h)
404 if i_level != self.num_resolutions - 1:
405 hs.append(self.down[i_level].downsample(hs[-1]))
406
407 # middle
408 h = hs[-1]
409 h = self.mid.block_1(h, temb)
410 h = self.mid.attn_1(h)
411 h = self.mid.block_2(h, temb)
412
413 # upsampling
414 for i_level in reversed(range(self.num_resolutions)):
415 for i_block in range(self.num_res_blocks + 1):
416 h = self.up[i_level].block[i_block](torch.cat([h, hs.pop()], dim=1), temb)
417 if len(self.up[i_level].attn) > 0:
418 h = self.up[i_level].attn[i_block](h)
419 if i_level != 0:
420 h = self.up[i_level].upsample(h)
421
422 # end
423 h = self.norm_out(h)
424 h = nonlinearity(h)
425 h = self.conv_out(h)
426 return h
427
428 def get_last_layer(self):
429 return self.conv_out.weight

Callers 1

forwardMethod · 0.45

Calls 3

appendMethod · 0.80
get_timestep_embeddingFunction · 0.70
nonlinearityFunction · 0.70

Tested by

no test coverage detected