(self, cfg: Config, *, parent: Optional[Module])
| 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. |
nothing calls this directly
no test coverage detected