(self,
input_feats=263,
latent_dim=256,
ff_size=1024,
num_layers=8,
num_heads=4,
dropout=0.1,
activation="gelu",
clip_dim=512,
clip_version=None,
guide_scale=1.0,
cond_mask_prob=0.1,
use_official_ckpt=False,
**kwargs)
| 36 | class MDMTransformer(nn.Module): |
| 37 | |
| 38 | def __init__(self, |
| 39 | input_feats=263, |
| 40 | latent_dim=256, |
| 41 | ff_size=1024, |
| 42 | num_layers=8, |
| 43 | num_heads=4, |
| 44 | dropout=0.1, |
| 45 | activation="gelu", |
| 46 | clip_dim=512, |
| 47 | clip_version=None, |
| 48 | guide_scale=1.0, |
| 49 | cond_mask_prob=0.1, |
| 50 | use_official_ckpt=False, |
| 51 | **kwargs): |
| 52 | super().__init__() |
| 53 | |
| 54 | self.latent_dim = latent_dim |
| 55 | self.ff_size = ff_size |
| 56 | self.num_layers = num_layers |
| 57 | self.num_heads = num_heads |
| 58 | self.dropout = dropout |
| 59 | self.activation = activation |
| 60 | self.clip_dim = clip_dim |
| 61 | self.input_feats = input_feats |
| 62 | self.guide_scale = guide_scale |
| 63 | self.use_official_ckpt = use_official_ckpt |
| 64 | |
| 65 | self.cond_mask_prob = cond_mask_prob |
| 66 | self.poseEmbedding = nn.Linear(self.input_feats, self.latent_dim) |
| 67 | self.sequence_pos_encoder = PositionalEncoding(self.latent_dim, |
| 68 | self.dropout) |
| 69 | |
| 70 | seqTransEncoderLayer = nn.TransformerEncoderLayer( |
| 71 | d_model=self.latent_dim, |
| 72 | nhead=self.num_heads, |
| 73 | dim_feedforward=self.ff_size, |
| 74 | dropout=self.dropout, |
| 75 | activation=self.activation) |
| 76 | |
| 77 | self.seqTransEncoder = nn.TransformerEncoder( |
| 78 | seqTransEncoderLayer, num_layers=self.num_layers) |
| 79 | |
| 80 | self.embed_timestep = TimestepEmbedder(self.latent_dim, |
| 81 | self.sequence_pos_encoder) |
| 82 | |
| 83 | self.embed_text = nn.Linear(self.clip_dim, self.latent_dim) |
| 84 | self.clip_version = clip_version |
| 85 | self.clip_model = self.load_and_freeze_clip(clip_version) |
| 86 | |
| 87 | self.poseFinal = nn.Linear(self.latent_dim, self.input_feats) |
| 88 | |
| 89 | def load_and_freeze_clip(self, clip_version): |
| 90 | clip_model, _ = clip.load(clip_version, device='cpu', jit=False) |
no test coverage detected