(self, global_pool=False)
| 24 | |
| 25 | class SDControlNet(torch.nn.Module): |
| 26 | def __init__(self, global_pool=False): |
| 27 | super().__init__() |
| 28 | self.time_proj = Timesteps(320) |
| 29 | self.time_embedding = torch.nn.Sequential( |
| 30 | torch.nn.Linear(320, 1280), |
| 31 | torch.nn.SiLU(), |
| 32 | torch.nn.Linear(1280, 1280) |
| 33 | ) |
| 34 | self.conv_in = torch.nn.Conv2d(4, 320, kernel_size=3, padding=1) |
| 35 | |
| 36 | self.controlnet_conv_in = ControlNetConditioningLayer(channels=(3, 16, 32, 96, 256, 320)) |
| 37 | |
| 38 | self.blocks = torch.nn.ModuleList([ |
| 39 | # CrossAttnDownBlock2D |
| 40 | ResnetBlock(320, 320, 1280), |
| 41 | AttentionBlock(8, 40, 320, 1, 768), |
| 42 | PushBlock(), |
| 43 | ResnetBlock(320, 320, 1280), |
| 44 | AttentionBlock(8, 40, 320, 1, 768), |
| 45 | PushBlock(), |
| 46 | DownSampler(320), |
| 47 | PushBlock(), |
| 48 | # CrossAttnDownBlock2D |
| 49 | ResnetBlock(320, 640, 1280), |
| 50 | AttentionBlock(8, 80, 640, 1, 768), |
| 51 | PushBlock(), |
| 52 | ResnetBlock(640, 640, 1280), |
| 53 | AttentionBlock(8, 80, 640, 1, 768), |
| 54 | PushBlock(), |
| 55 | DownSampler(640), |
| 56 | PushBlock(), |
| 57 | # CrossAttnDownBlock2D |
| 58 | ResnetBlock(640, 1280, 1280), |
| 59 | AttentionBlock(8, 160, 1280, 1, 768), |
| 60 | PushBlock(), |
| 61 | ResnetBlock(1280, 1280, 1280), |
| 62 | AttentionBlock(8, 160, 1280, 1, 768), |
| 63 | PushBlock(), |
| 64 | DownSampler(1280), |
| 65 | PushBlock(), |
| 66 | # DownBlock2D |
| 67 | ResnetBlock(1280, 1280, 1280), |
| 68 | PushBlock(), |
| 69 | ResnetBlock(1280, 1280, 1280), |
| 70 | PushBlock(), |
| 71 | # UNetMidBlock2DCrossAttn |
| 72 | ResnetBlock(1280, 1280, 1280), |
| 73 | AttentionBlock(8, 160, 1280, 1, 768), |
| 74 | ResnetBlock(1280, 1280, 1280), |
| 75 | PushBlock() |
| 76 | ]) |
| 77 | |
| 78 | self.controlnet_blocks = torch.nn.ModuleList([ |
| 79 | torch.nn.Conv2d(320, 320, kernel_size=(1, 1)), |
| 80 | torch.nn.Conv2d(320, 320, kernel_size=(1, 1), bias=False), |
| 81 | torch.nn.Conv2d(320, 320, kernel_size=(1, 1), bias=False), |
| 82 | torch.nn.Conv2d(320, 320, kernel_size=(1, 1), bias=False), |
| 83 | torch.nn.Conv2d(640, 640, kernel_size=(1, 1)), |
no test coverage detected