HunYuan config fix
This commit is contained in:
@@ -1,7 +1,7 @@
|
|||||||
{
|
{
|
||||||
"_name_or_path": "mt5",
|
"_name_or_path": "mt5",
|
||||||
"architectures": [
|
"architectures": [
|
||||||
"MT5ForConditionalGeneration"
|
"MT5EncoderModel"
|
||||||
],
|
],
|
||||||
"classifier_dropout": 0.0,
|
"classifier_dropout": 0.0,
|
||||||
"d_ff": 5120,
|
"d_ff": 5120,
|
||||||
|
|||||||
@@ -77,6 +77,5 @@ def load_hydit(model_path, model_conf):
|
|||||||
load_device = load_device,
|
load_device = load_device,
|
||||||
offload_device = offload_device,
|
offload_device = offload_device,
|
||||||
current_device = "cpu",
|
current_device = "cpu",
|
||||||
size = 6 * (1024**3),
|
|
||||||
)
|
)
|
||||||
return model_patcher
|
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.final_layer = FinalLayer(hidden_size, hidden_size, patch_size, self.out_channels)
|
||||||
self.unpatchify_channels = self.out_channels
|
self.unpatchify_channels = self.out_channels
|
||||||
|
|
||||||
# probably not needed when not training?
|
|
||||||
# self.initialize_weights()
|
|
||||||
|
|
||||||
def forward_raw(self,
|
def forward_raw(self,
|
||||||
x,
|
x,
|
||||||
t,
|
t,
|
||||||
@@ -415,40 +412,6 @@ class HunYuanDiT(nn.Module):
|
|||||||
else:
|
else:
|
||||||
return out
|
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):
|
def unpatchify(self, x, h, w):
|
||||||
"""
|
"""
|
||||||
x: (N, T, patch_size**2 * C)
|
x: (N, T, patch_size**2 * C)
|
||||||
|
|||||||
Reference in New Issue
Block a user