MCPcopy Create free account
hub / github.com/apple/axlearn / __init__

Method __init__

axlearn/common/dit.py:159–168  ·  view source on GitHub ↗
(self, cfg: Config, *, parent: Optional[Module])

Source from the content-addressed store, hash-verified

157 emb: Embedding.Config = Embedding.default_config()
158
159 def __init__(self, cfg: Config, *, parent: Optional[Module]):
160 super().__init__(cfg, parent=parent)
161 cfg = self.config
162 # Ref: https://github.com/facebookresearch/DiT/blob/main/models.py#L74
163 # The num_embeddings = cfg.num_classes + (dropout_rate > 0)
164 # Additional token is added to label dropout for classifier-free guidance.
165 use_drop_rate = int(cfg.dropout_rate > 0)
166 self._add_child(
167 "emb", cfg.emb.set(dim=cfg.output_dim, num_embeddings=cfg.num_classes + use_drop_rate)
168 )
169
170 def _mask_label(self, label: Tensor) -> Tensor:
171 """This function is activated only when dropout_rate > 0 and is_training=False.

Callers

nothing calls this directly

Calls 3

_add_childMethod · 0.80
__init__Method · 0.45
setMethod · 0.45

Tested by

no test coverage detected