Fix SteadyDancer GGUF dtypes
This commit is contained in:
@@ -53,21 +53,25 @@ class DYModule(nn.Module):
|
||||
self.bn2 = nn.GroupNorm(num_groups=4, num_channels=self.dim)
|
||||
|
||||
def forward(self, x):
|
||||
r = self.conv(x)
|
||||
|
||||
x_type = x.dtype
|
||||
r = self.conv(x.to(self.conv.weight.dtype)).to(x_type)
|
||||
b, c, h, w = x.size()
|
||||
|
||||
y = self.avg_pool(x).view(b, c * self.mul)
|
||||
y = self.fc(y)
|
||||
dy_phi = self.fc_phi(y).view(b, self.dim, self.dim)
|
||||
dy_scale = self.hs(self.fc_scale(y)).view(b, -1, 1, 1)
|
||||
r = dy_scale.expand_as(r) * r
|
||||
|
||||
x = self.conv_q(x)
|
||||
x = self.conv_q(x.to(self.conv_q.weight.dtype)).to(self.bn1.weight.dtype)
|
||||
x = self.bn1(x)
|
||||
|
||||
x = x.view(b, -1, h * w)
|
||||
x = self.bn2(torch.matmul(dy_phi, x)) + x
|
||||
x = x + self.bn2(torch.matmul(dy_phi, x.to(dy_phi.dtype)).to(self.bn2.weight.dtype))
|
||||
x = x.view(b, -1, h, w)
|
||||
x = self.conv_p(x)
|
||||
|
||||
x = self.conv_p(x.to(self.conv_p.weight.dtype)).to(x_type)
|
||||
|
||||
return x + r
|
||||
|
||||
|
||||
|
||||
+12
-17
@@ -44,11 +44,10 @@ class FactorConv3d(nn.Module):
|
||||
self.act = nn.SiLU()
|
||||
|
||||
def forward(self, x):
|
||||
x = self.spatial(x)
|
||||
out_dtype = x.dtype
|
||||
x = self.spatial(x.to(self.spatial.weight.dtype)).to(out_dtype)
|
||||
x = self.act(x)
|
||||
x = self.temporal(x)
|
||||
return x
|
||||
|
||||
return self.temporal(x.to(self.temporal.weight.dtype)).to(out_dtype)
|
||||
|
||||
class LayerNorm2D(nn.Module):
|
||||
"""
|
||||
@@ -110,29 +109,25 @@ class PoseRefNetNoBNV3(nn.Module):
|
||||
return: (B, d_model, T, H, W)
|
||||
"""
|
||||
B, _, T, H, W = pose.shape
|
||||
L = H * W
|
||||
|
||||
p_trans = pose.permute(0, 2, 1, 3, 4).contiguous().flatten(0, 1)
|
||||
r_trans = ref.permute(0, 2, 1, 3, 4).contiguous().flatten(0, 1)
|
||||
|
||||
p_trans = self.proj_p(p_trans)
|
||||
r_trans = self.proj_r(r_trans)
|
||||
|
||||
p_trans = p_trans.flatten(2).transpose(1, 2)
|
||||
r_trans = r_trans.flatten(2).transpose(1, 2)
|
||||
p_trans = self.proj_p(p_trans.to(self.proj_p.weight.dtype)).to(self.cross_attn.in_proj_weight.dtype).flatten(2).transpose(1, 2)
|
||||
r_trans = self.proj_r(r_trans.to(self.proj_r.weight.dtype)).to(self.cross_attn.in_proj_weight.dtype).flatten(2).transpose(1, 2)
|
||||
|
||||
out = self.cross_attn(query=r_trans,
|
||||
key=p_trans,
|
||||
value=p_trans,
|
||||
key_padding_mask=mask)[0]
|
||||
|
||||
out = out.transpose(1, 2).contiguous().view(B*T, -1, H, W)
|
||||
out = self.norm1(out)
|
||||
out = self.norm1(out.transpose(1, 2).contiguous().view(B*T, -1, H, W))
|
||||
|
||||
ffn_out = self.ffn_pose(out)
|
||||
out = out + ffn_out
|
||||
out_type = out.dtype
|
||||
|
||||
out = out + self.ffn_pose(out.to(self.ffn_pose[0].weight.dtype)).to(out_type)
|
||||
out = self.norm2(out)
|
||||
out = self.proj_p_back(out)
|
||||
out = out.view(B, T, -1, H, W).contiguous().transpose(1, 2)
|
||||
|
||||
return out
|
||||
out = self.proj_p_back(out.to(self.proj_p_back.weight.dtype)).to(out_type)
|
||||
|
||||
return out.view(B, T, -1, H, W).contiguous().transpose(1, 2)
|
||||
|
||||
Reference in New Issue
Block a user