HunYuan config fix
This commit is contained in:
@@ -1,7 +1,7 @@
|
||||
{
|
||||
"_name_or_path": "mt5",
|
||||
"architectures": [
|
||||
"MT5ForConditionalGeneration"
|
||||
"MT5EncoderModel"
|
||||
],
|
||||
"classifier_dropout": 0.0,
|
||||
"d_ff": 5120,
|
||||
|
||||
@@ -77,6 +77,5 @@ def load_hydit(model_path, model_conf):
|
||||
load_device = load_device,
|
||||
offload_device = offload_device,
|
||||
current_device = "cpu",
|
||||
size = 6 * (1024**3),
|
||||
)
|
||||
return model_patcher
|
||||
|
||||
@@ -244,9 +244,6 @@ class HunYuanDiT(nn.Module):
|
||||
self.final_layer = FinalLayer(hidden_size, hidden_size, patch_size, self.out_channels)
|
||||
self.unpatchify_channels = self.out_channels
|
||||
|
||||
# probably not needed when not training?
|
||||
# self.initialize_weights()
|
||||
|
||||
def forward_raw(self,
|
||||
x,
|
||||
t,
|
||||
@@ -415,40 +412,6 @@ class HunYuanDiT(nn.Module):
|
||||
else:
|
||||
return out
|
||||
|
||||
|
||||
def initialize_weights(self):
|
||||
# Initialize transformer layers:
|
||||
def _basic_init(module):
|
||||
if isinstance(module, nn.Linear):
|
||||
torch.nn.init.xavier_uniform_(module.weight)
|
||||
if module.bias is not None:
|
||||
nn.init.constant_(module.bias, 0)
|
||||
self.apply(_basic_init)
|
||||
|
||||
# Initialize patch_embed like nn.Linear (instead of nn.Conv2d):
|
||||
w = self.x_embedder.proj.weight.data
|
||||
nn.init.xavier_uniform_(w.view([w.shape[0], -1]))
|
||||
nn.init.constant_(self.x_embedder.proj.bias, 0)
|
||||
|
||||
# Initialize label embedding table:
|
||||
nn.init.normal_(self.extra_embedder[0].weight, std=0.02)
|
||||
nn.init.normal_(self.extra_embedder[2].weight, std=0.02)
|
||||
|
||||
# Initialize timestep embedding MLP:
|
||||
nn.init.normal_(self.t_embedder.mlp[0].weight, std=0.02)
|
||||
nn.init.normal_(self.t_embedder.mlp[2].weight, std=0.02)
|
||||
|
||||
# Zero-out adaLN modulation layers in HunYuanDiT blocks:
|
||||
for block in self.blocks:
|
||||
nn.init.constant_(block.default_modulation[-1].weight, 0)
|
||||
nn.init.constant_(block.default_modulation[-1].bias, 0)
|
||||
|
||||
# Zero-out output layers:
|
||||
nn.init.constant_(self.final_layer.adaLN_modulation[-1].weight, 0)
|
||||
nn.init.constant_(self.final_layer.adaLN_modulation[-1].bias, 0)
|
||||
nn.init.constant_(self.final_layer.linear.weight, 0)
|
||||
nn.init.constant_(self.final_layer.linear.bias, 0)
|
||||
|
||||
def unpatchify(self, x, h, w):
|
||||
"""
|
||||
x: (N, T, patch_size**2 * C)
|
||||
|
||||
Reference in New Issue
Block a user