(self, config: PretrainedConfig)
| 201 | class DiT(PretrainedModel): |
| 202 | |
| 203 | def __init__(self, config: PretrainedConfig): |
| 204 | self.check_config(config) |
| 205 | super().__init__(config) |
| 206 | self.learn_sigma = config.learn_sigma |
| 207 | self.in_channels = config.in_channels |
| 208 | self.out_channels = config.in_channels * 2 if config.learn_sigma else config.in_channels |
| 209 | self.input_size = config.input_size |
| 210 | self.patch_size = config.patch_size |
| 211 | self.num_heads = config.num_attention_heads |
| 212 | self.dtype = str_dtype_to_trt(config.dtype) |
| 213 | self.cfg_scale = config.cfg_scale |
| 214 | self.mapping = config.mapping |
| 215 | |
| 216 | self.x_embedder = PatchEmbed(config.input_size, |
| 217 | config.patch_size, |
| 218 | config.in_channels, |
| 219 | config.hidden_size, |
| 220 | bias=True, |
| 221 | dtype=self.dtype) |
| 222 | self.t_embedder = TimestepEmbedder(config.hidden_size, dtype=self.dtype) |
| 223 | self.y_embedder = LabelEmbedder(config.num_classes, |
| 224 | config.hidden_size, |
| 225 | config.class_dropout_prob, |
| 226 | dtype=self.dtype) |
| 227 | num_patches = self.x_embedder.num_patches |
| 228 | |
| 229 | self.pos_embed = Parameter(shape=(1, num_patches, config.hidden_size), |
| 230 | dtype=self.dtype) |
| 231 | self.blocks = ModuleList([ |
| 232 | DiTBlock(config.hidden_size, |
| 233 | config.num_attention_heads, |
| 234 | mlp_ratio=config.mlp_ratio, |
| 235 | mapping=config.mapping, |
| 236 | dtype=self.dtype, |
| 237 | quant_mode=config.quant_mode) |
| 238 | for _ in range(config.num_hidden_layers) |
| 239 | ]) |
| 240 | self.final_layer = FinalLayer(config.hidden_size, |
| 241 | config.patch_size, |
| 242 | self.out_channels, |
| 243 | mapping=config.mapping, |
| 244 | dtype=self.dtype) |
| 245 | |
| 246 | # We need to invoke default `__post_init__()` for quantized layers. |
| 247 | # def __post_init__(self): |
no test coverage detected