MCPcopy Create free account
hub / github.com/CompVis/zigma / LabelEmbedder

Class LabelEmbedder

model_zigma.py:278–310  ·  view source on GitHub ↗

Embeds class labels into vector representations. Also handles label dropout for classifier-free guidance.

Source from the content-addressed store, hash-verified

276
277
278class LabelEmbedder(nn.Module):
279 """
280 Embeds class labels into vector representations. Also handles label dropout for classifier-free guidance.
281 """
282
283 def __init__(self, num_classes, hidden_size, dropout_prob):
284 super().__init__()
285 use_cfg_embedding = dropout_prob > 0
286 self.embedding_table = nn.Embedding(
287 num_classes + use_cfg_embedding, hidden_size
288 )
289 self.num_classes = num_classes
290 self.dropout_prob = dropout_prob
291
292 def token_drop(self, labels, force_drop_ids=None):
293 """
294 Drops labels to enable classifier-free guidance.
295 """
296 if force_drop_ids is None:
297 drop_ids = (
298 torch.rand(labels.shape[0], device=labels.device) < self.dropout_prob
299 )
300 else:
301 drop_ids = force_drop_ids == 1
302 labels = torch.where(drop_ids, self.num_classes, labels)
303 return labels
304
305 def forward(self, labels, train, force_drop_ids=None):
306 use_dropout = self.dropout_prob > 0
307 if (train and use_dropout) or (force_drop_ids is not None):
308 labels = self.token_drop(labels, force_drop_ids)
309 embeddings = self.embedding_table(labels)
310 return embeddings
311
312
313class FinalLayer(nn.Module):

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected