File size: 4,465 Bytes
d4cbafd | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 | 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
'''
# print(h.shape)
h_feat = self.encode_past(h.reshape(h.size(0), -1)).unsqueeze(1)
# print(h_feat.shape)
# n_samples, 1, 64
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) # (B, 1, 1)
mask = mask.float().masked_fill(mask == 0, float('-inf')).masked_fill(mask == 1, float(0.0))
context = self.encoder_context(context, mask)
# context = context.view(batch_size, 1, -1) # (B, 1, F)
time_emb = torch.cat([beta, torch.sin(beta), torch.cos(beta)], dim=-1) # (B, 1, 3)
ctx_emb = torch.cat([time_emb, context], dim=-1) # (B, 1, F+3)
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) # (B, 1, 1)
mask = mask.float().masked_fill(mask == 0, float('-inf')).masked_fill(mask == 1, float(0.0))
context = self.encoder_context(context, mask)
# context = context.view(batch_size, 1, -1) # (B, 1, F)
time_emb = torch.cat([beta, torch.sin(beta), torch.cos(beta)], dim=-1) # (B, 1, 3)
# time_emb: [11, 1, 3]
# context: [11, 1, 256]
ctx_emb = torch.cat([time_emb, context], dim=-1).repeat(1, 10, 1).unsqueeze(2)
# x: 11, 10, 20, 2
# ctx_emb: 11, 10, 1, 259
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: 11, 10, 20, 512
trans = self.concat3.batch_generate(ctx_emb, trans)
trans = self.concat4.batch_generate(ctx_emb, trans)
return self.linear.batch_generate(ctx_emb, trans)
|