(name, ckpt_path)
| 44 | return model |
| 45 | |
| 46 | def get_text_model(name, ckpt_path): |
| 47 | if name == 'kit_ml': |
| 48 | model = build_submodule(dict( |
| 49 | type='T2MTextEncoder', |
| 50 | word_size=300, |
| 51 | pos_size=15, |
| 52 | hidden_size=512, |
| 53 | output_size=512, |
| 54 | max_text_len=20 |
| 55 | )) |
| 56 | elif name == 'kit_ttc' and 'human' not in ckpt_path: |
| 57 | model = build_submodule(dict( |
| 58 | type='TextEncoder', |
| 59 | pretrained_model='clip', |
| 60 | text_latent_dim=256, |
| 61 | time_embed_dim=512, |
| 62 | dropout=0, |
| 63 | num_text_layers=2, |
| 64 | text_num_heads=4, |
| 65 | text_ff_size=2048, |
| 66 | use_text_proj=True |
| 67 | )) |
| 68 | elif name == 'kit_ttc': |
| 69 | model = build_submodule(dict( |
| 70 | type='TextEncoder', |
| 71 | pretrained_model='clip', |
| 72 | text_latent_dim=256, |
| 73 | time_embed_dim=512, |
| 74 | dropout=0, |
| 75 | num_text_layers=4, |
| 76 | text_num_heads=4, |
| 77 | text_ff_size=2048, |
| 78 | use_text_proj=True |
| 79 | )) |
| 80 | else: |
| 81 | model = build_submodule(dict( |
| 82 | type='T2MTextEncoder', |
| 83 | word_size=300, |
| 84 | pos_size=15, |
| 85 | hidden_size=512, |
| 86 | output_size=512, |
| 87 | max_text_len=20 |
| 88 | )) |
| 89 | model.load_pretrained(ckpt_path) |
| 90 | return model |
no test coverage detected