| 23 | |
| 24 | |
| 25 | class ClassEmbedder(nn.Module): |
| 26 | def __init__(self, embed_dim, n_classes=1000, key='class', ucg_rate=0.1): |
| 27 | super().__init__() |
| 28 | self.key = key |
| 29 | self.embedding = nn.Embedding(n_classes, embed_dim) |
| 30 | self.n_classes = n_classes |
| 31 | self.ucg_rate = ucg_rate |
| 32 | |
| 33 | def forward(self, batch, key=None, disable_dropout=False): |
| 34 | if key is None: |
| 35 | key = self.key |
| 36 | # this is for use in crossattn |
| 37 | c = batch[key][:, None] |
| 38 | if self.ucg_rate > 0. and not disable_dropout: |
| 39 | mask = 1. - torch.bernoulli(torch.ones_like(c) * self.ucg_rate) |
| 40 | c = mask * c + (1-mask) * torch.ones_like(c)*(self.n_classes-1) |
| 41 | c = c.long() |
| 42 | c = self.embedding(c) |
| 43 | return c |
| 44 | |
| 45 | def get_unconditional_conditioning(self, bs, device="cuda"): |
| 46 | uc_class = self.n_classes - 1 # 1000 classes --> 0 ... 999, one extra class for ucg (class 1000) |
| 47 | uc = torch.ones((bs,), device=device) * uc_class |
| 48 | uc = {self.key: uc} |
| 49 | return uc |
| 50 | |
| 51 | |
| 52 | def disabled_train(self, mode=True): |
nothing calls this directly
no outgoing calls
no test coverage detected