| import math |
| import torch |
| import torch.nn as nn |
| from torch.nn import Module, Linear |
|
|
| from models.layers import PositionalEncoding, ConcatSquashLinear |
|
|
| class st_encoder(nn.Module): |
| def __init__(self): |
| super().__init__() |
| channel_in = 2 |
| channel_out = 32 |
| dim_kernel = 3 |
| self.dim_embedding_key = 256 |
| self.spatial_conv = nn.Conv1d(channel_in, channel_out, dim_kernel, stride=1, padding=1) |
| self.temporal_encoder = nn.GRU(channel_out, self.dim_embedding_key, 1, batch_first=True) |
|
|
| self.relu = nn.ReLU() |
|
|
| self.reset_parameters() |
|
|
| def reset_parameters(self): |
| nn.init.kaiming_normal_(self.spatial_conv.weight) |
| nn.init.kaiming_normal_(self.temporal_encoder.weight_ih_l0) |
| nn.init.kaiming_normal_(self.temporal_encoder.weight_hh_l0) |
| nn.init.zeros_(self.spatial_conv.bias) |
| nn.init.zeros_(self.temporal_encoder.bias_ih_l0) |
| nn.init.zeros_(self.temporal_encoder.bias_hh_l0) |
|
|
| def forward(self, X): |
| ''' |
| X: b, T, 2 |
| |
| return: b, F |
| ''' |
| X_t = torch.transpose(X, 1, 2) |
| X_after_spatial = self.relu(self.spatial_conv(X_t)) |
| X_embed = torch.transpose(X_after_spatial, 1, 2) |
|
|
| output_x, state_x = self.temporal_encoder(X_embed) |
| state_x = state_x.squeeze(0) |
|
|
| return state_x |
|
|
|
|
| class social_transformer(nn.Module): |
| def __init__(self, past_len=10, in_channels=6): |
| super(social_transformer, self).__init__() |
| self.encode_past = nn.Linear(past_len * in_channels, 256, bias=False) |
| self.layer = nn.TransformerEncoderLayer(d_model=256, nhead=2, dim_feedforward=256) |
| self.transformer_encoder = nn.TransformerEncoder(self.layer, num_layers=2) |
|
|
| def forward(self, h, mask): |
| ''' |
| h: batch_size, t, 2 |
| ''' |
| |
| h_feat = self.encode_past(h.reshape(h.size(0), -1)).unsqueeze(1) |
| |
| |
| h_feat_ = self.transformer_encoder(h_feat, mask) |
| h_feat = h_feat + h_feat_ |
|
|
| return h_feat |
|
|
|
|
| class TransformerDenoisingModel(Module): |
|
|
| def __init__(self, context_dim=256, tf_layer=2, past_len=10): |
| super().__init__() |
| self.encoder_context = social_transformer(past_len=past_len) |
| self.pos_emb = PositionalEncoding(d_model=2*context_dim, dropout=0.1, max_len=24) |
| self.concat1 = ConcatSquashLinear(2, 2*context_dim, context_dim+3) |
| self.layer = nn.TransformerEncoderLayer(d_model=2*context_dim, nhead=2, dim_feedforward=2*context_dim) |
| self.transformer_encoder = nn.TransformerEncoder(self.layer, num_layers=tf_layer) |
| self.concat3 = ConcatSquashLinear(2*context_dim,context_dim,context_dim+3) |
| self.concat4 = ConcatSquashLinear(context_dim,context_dim//2,context_dim+3) |
| self.linear = ConcatSquashLinear(context_dim//2, 2, context_dim+3) |
|
|
|
|
| def forward(self, x, beta, context, mask): |
| batch_size = x.size(0) |
| beta = beta.view(batch_size, 1, 1) |
| mask = mask.float().masked_fill(mask == 0, float('-inf')).masked_fill(mask == 1, float(0.0)) |
| context = self.encoder_context(context, mask) |
| |
|
|
| time_emb = torch.cat([beta, torch.sin(beta), torch.cos(beta)], dim=-1) |
| ctx_emb = torch.cat([time_emb, context], dim=-1) |
| |
| x = self.concat1(ctx_emb, x) |
| final_emb = x.permute(1,0,2) |
| final_emb = self.pos_emb(final_emb) |
| |
| trans = self.transformer_encoder(final_emb).permute(1,0,2) |
| trans = self.concat3(ctx_emb, trans) |
| trans = self.concat4(ctx_emb, trans) |
| return self.linear(ctx_emb, trans) |
| |
|
|
| def generate_accelerate(self, x, beta, context, mask): |
| batch_size = x.size(0) |
| beta = beta.view(beta.size(0), 1, 1) |
| mask = mask.float().masked_fill(mask == 0, float('-inf')).masked_fill(mask == 1, float(0.0)) |
| context = self.encoder_context(context, mask) |
| |
|
|
| time_emb = torch.cat([beta, torch.sin(beta), torch.cos(beta)], dim=-1) |
| |
| |
| ctx_emb = torch.cat([time_emb, context], dim=-1).repeat(1, 10, 1).unsqueeze(2) |
| |
| |
| K = x.size(1) |
| T = x.size(2) |
| D = 2 * 256 |
| x = self.concat1.batch_generate(ctx_emb, x).contiguous().view(-1, T, D) |
| final_emb = x.permute(1, 0, 2) |
| final_emb = self.pos_emb(final_emb) |
|
|
| trans = self.transformer_encoder(final_emb).permute(1, 0, 2).contiguous().view(-1, K, T, D) |
| |
| trans = self.concat3.batch_generate(ctx_emb, trans) |
| trans = self.concat4.batch_generate(ctx_emb, trans) |
| return self.linear.batch_generate(ctx_emb, trans) |
|
|