(self,
num_retrieval=None,
topk=None,
retrieval_file=None,
latent_dim=512,
output_dim=512,
num_layers=2,
num_motion_layers=4,
kinematic_coef=0.1,
max_seq_len=196,
num_heads=8,
ff_size=1024,
stride=4,
sa_block_cfg=None,
ffn_cfg=None,
dropout=0)
| 46 | class RetrievalDatabase(nn.Module): |
| 47 | |
| 48 | def __init__(self, |
| 49 | num_retrieval=None, |
| 50 | topk=None, |
| 51 | retrieval_file=None, |
| 52 | latent_dim=512, |
| 53 | output_dim=512, |
| 54 | num_layers=2, |
| 55 | num_motion_layers=4, |
| 56 | kinematic_coef=0.1, |
| 57 | max_seq_len=196, |
| 58 | num_heads=8, |
| 59 | ff_size=1024, |
| 60 | stride=4, |
| 61 | sa_block_cfg=None, |
| 62 | ffn_cfg=None, |
| 63 | dropout=0): |
| 64 | super().__init__() |
| 65 | self.num_retrieval = num_retrieval |
| 66 | self.topk = topk |
| 67 | self.latent_dim = latent_dim |
| 68 | self.stride = stride |
| 69 | self.kinematic_coef = kinematic_coef |
| 70 | self.num_layers = num_layers |
| 71 | self.num_motion_layers = num_motion_layers |
| 72 | self.max_seq_len = max_seq_len |
| 73 | data = np.load(retrieval_file) |
| 74 | self.text_features = torch.Tensor(data['text_features']) |
| 75 | self.captions = data['captions'] |
| 76 | self.motions = data['motions'] |
| 77 | self.m_lengths = data['m_lengths'] |
| 78 | self.clip_seq_features = data['clip_seq_features'] |
| 79 | self.train_indexes = data.get('train_indexes', None) |
| 80 | self.test_indexes = data.get('test_indexes', None) |
| 81 | |
| 82 | self.latent_dim = latent_dim |
| 83 | self.output_dim = output_dim |
| 84 | self.motion_proj = nn.Linear(self.motions.shape[-1], self.latent_dim) |
| 85 | self.motion_pos_embedding = nn.Parameter( |
| 86 | torch.randn(max_seq_len, self.latent_dim)) |
| 87 | self.motion_encoder_blocks = nn.ModuleList() |
| 88 | for i in range(num_motion_layers): |
| 89 | self.motion_encoder_blocks.append( |
| 90 | EncoderLayer(sa_block_cfg=sa_block_cfg, ffn_cfg=ffn_cfg)) |
| 91 | TransEncoderLayer = nn.TransformerEncoderLayer(d_model=self.latent_dim, |
| 92 | nhead=num_heads, |
| 93 | dim_feedforward=ff_size, |
| 94 | dropout=dropout, |
| 95 | activation="gelu") |
| 96 | self.text_encoder = nn.TransformerEncoder(TransEncoderLayer, |
| 97 | num_layers=num_layers) |
| 98 | self.results = {} |
| 99 | |
| 100 | def extract_text_feature(self, text, clip_model, device): |
| 101 | text = clip.tokenize([text], truncate=True).to(device) |
nothing calls this directly
no test coverage detected