MCPcopy Create free account
hub / github.com/Alpha-VLLM/LLaMA2-Accessory / ParallelLabelEmbedder

Class ParallelLabelEmbedder

Large-DiT-ImageNet/models.py:87–126  ·  view source on GitHub ↗

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

Source from the content-addressed store, hash-verified

85
86
87class ParallelLabelEmbedder(nn.Module):
88 r"""Embeds class labels into vector representations. Also handles label
89 dropout for classifier-free guidance.
90 """
91 def __init__(self, num_classes, hidden_size, dropout_prob):
92 super().__init__()
93 use_cfg_embedding = int(dropout_prob > 0)
94 self.embedding_table = ParallelEmbedding(
95 num_classes + use_cfg_embedding, hidden_size,
96 init_method=functools.partial(nn.init.normal_, std=0.02),
97 )
98 self.num_classes = num_classes
99 self.dropout_prob = dropout_prob
100
101 def token_drop(self, labels, force_drop_ids=None):
102 """
103 Drops labels to enable classifier-free guidance.
104 """
105 if force_drop_ids is None:
106 drop_ids = torch.rand(
107 labels.shape[0], device=labels.device
108 ) < self.dropout_prob
109 drop_ids = drop_ids.cuda()
110 dist.broadcast(
111 drop_ids,
112 fs_init.get_model_parallel_src_rank(),
113 fs_init.get_model_parallel_group(),
114 )
115 drop_ids = drop_ids.to(labels.device)
116 else:
117 drop_ids = force_drop_ids == 1
118 labels = torch.where(drop_ids, self.num_classes, labels)
119 return labels
120
121 def forward(self, labels, train, force_drop_ids=None):
122 use_dropout = self.dropout_prob > 0
123 if (train and use_dropout) or (force_drop_ids is not None):
124 labels = self.token_drop(labels, force_drop_ids)
125 embeddings = self.embedding_table(labels)
126 return embeddings
127
128
129#############################################################################

Callers 1

__init__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected