Make mlp args match the one from timm

This commit is contained in:
City
2024-12-11 22:37:13 +01:00
parent d1f751677a
commit ca2bc8bc65
+8 -3
View File
@@ -419,20 +419,25 @@ class LabelEmbedder(nn.Module):
embeddings = self.embedding_table(labels) embeddings = self.embedding_table(labels)
return embeddings return embeddings
class Mlp(nn.Module): class Mlp(nn.Module):
def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, dtype=None, device=None, operations=None) -> None: def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, bias=True, drop=None, dtype=None, device=None, operations=None) -> None:
super().__init__() super().__init__()
out_features = out_features or in_features out_features = out_features or in_features
hidden_features = hidden_features or in_features hidden_features = hidden_features or in_features
self.fc1 = operations.Linear(in_features, hidden_features, bias=True, dtype=dtype, device=device) self.fc1 = operations.Linear(in_features, hidden_features, bias=bias, dtype=dtype, device=device)
self.act = act_layer() self.act = act_layer()
self.fc2 = operations.Linear(hidden_features, out_features, bias=True, dtype=dtype, device=device) self.fc2 = operations.Linear(hidden_features, out_features, bias=bias, dtype=dtype, device=device)
self.drop1 = nn.Identity()
self.drop2 = nn.Identity()
def forward(self, x: torch.Tensor) -> torch.Tensor: def forward(self, x: torch.Tensor) -> torch.Tensor:
x = self.act(self.fc1(x)) x = self.act(self.fc1(x))
return self.fc2(x) return self.fc2(x)
class CaptionEmbedder(nn.Module): class CaptionEmbedder(nn.Module):
""" """
Embeds class labels into vector representations. Also handles label dropout for classifier-free guidance. Embeds class labels into vector representations. Also handles label dropout for classifier-free guidance.