Spaces:
Running
Running
File size: 6,499 Bytes
17f1f54 | 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 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 | # # coding: utf-8
import torch
import torch.nn as nn
import math
from torch import Tensor
from helpers import freeze_params, subsequent_mask
from transformer_layers import PositionalEncoding, TransformerDecoderLayer
class SinusoidalPositionEmbeddings(nn.Module):
def __init__(self, dim: int):
super().__init__()
self.dim = dim
def forward(self, time: Tensor) -> Tensor:
# time: [B] (long or float)
device = time.device
half_dim = self.dim // 2
freq = math.log(10000) / (half_dim - 1)
freq = torch.exp(torch.arange(half_dim, device=device) * -freq)
# ensure float
time = time.float()
# [B, half_dim]
angles = time[:, None] * freq[None, :]
# [B, dim]
return torch.cat((angles.sin(), angles.cos()), dim=-1)
class ACD_Denoiser(nn.Module):
def __init__(
self,
num_layers: int = 2,
num_heads: int = 4,
hidden_size: int = 512,
ff_size: int = 2048,
dropout: float = 0.1,
emb_dropout: float = 0.1,
vocab_size: int = 1,
freeze: bool = False,
trg_size: int = 150,
decoder_trg_trg_: bool = True,
**kwargs
):
super(ACD_Denoiser, self).__init__()
# remember for repr
self.num_layers = num_layers
self.num_heads = num_heads
# Input features = joints (trg_size=150) + iconicity/bone dir+len (50*4)
# total in_feature_size = 150 + 200 = 350 (= 50 * 7)
self.in_feature_size = trg_size + (trg_size // 3) * 4
self.out_feature_size = trg_size
# Embedding for target features
self.pos_drop = nn.Dropout(p=emb_dropout)
self.trg_embed = nn.Linear(self.in_feature_size, hidden_size)
self.pe = PositionalEncoding(hidden_size, mask_count=True)
self.emb_dropout = nn.Dropout(p=emb_dropout)
# Two-layer decoder stack (as in original)
if num_layers == 2:
self.layers_pose_condition = TransformerDecoderLayer(
size=hidden_size,
ff_size=ff_size,
num_heads=num_heads,
dropout=dropout,
decoder_trg_trg=decoder_trg_trg_,
)
self.layer_norm_mid = nn.LayerNorm(hidden_size, eps=1e-6)
self.output_layer_mid = nn.Linear(hidden_size, self.in_feature_size, bias=False)
self.o1_embed = nn.Linear(trg_size, hidden_size) # joints part (50*3)
self.o2_embed = nn.Linear((trg_size // 3) * 4, hidden_size) # bones part (50*4)
self.layers_mha_ac = TransformerDecoderLayer(
size=hidden_size,
ff_size=ff_size,
num_heads=num_heads,
dropout=dropout,
decoder_trg_trg=decoder_trg_trg_,
)
self.layer_norm = nn.LayerNorm(hidden_size, eps=1e-6)
# --- time embedding ---
self.time_mlp = nn.Sequential(
SinusoidalPositionEmbeddings(hidden_size),
nn.Linear(hidden_size, hidden_size * 2),
nn.GELU(),
nn.Linear(hidden_size * 2, hidden_size),
)
# NEW: small projector to inject [sigma_B, sigma_H] (2 scalars) into the time embedding
self.time_proj = nn.Sequential(
nn.Linear(hidden_size + 2, hidden_size),
nn.GELU(),
nn.Linear(hidden_size, hidden_size),
)
# Output head -> predict x0 joints (trg_size)
self.output_layer = nn.Linear(hidden_size, trg_size, bias=False)
if freeze:
freeze_params(self)
def forward(
self,
t: Tensor,
trg_embed: Tensor = None,
encoder_output: Tensor = None,
src_mask: Tensor = None,
trg_mask: Tensor = None,
sigma_B: Tensor = None, # NEW (optional): [B]
sigma_H: Tensor = None, # NEW (optional): [B]
**kwargs,
) -> Tensor:
assert trg_mask is not None, "trg_mask required for Transformer"
# --- time conditioning ---
# base time embedding: [B, hidden]
t_base = self.time_mlp(t)
# add two-rate noise indicators; default to zeros for backward-compat
if sigma_B is None or sigma_H is None:
# type/shape safety
sigma_B = torch.zeros_like(t, dtype=t_base.dtype)
sigma_H = torch.zeros_like(t, dtype=t_base.dtype)
# concat and project back to hidden
t_aug = torch.stack([sigma_B, sigma_H], dim=-1) # [B, 2]
t_cond = self.time_proj(torch.cat([t_base, t_aug], dim=-1)) # [B, hidden]
# broadcast over time dimension of encoder_output
time_embed = t_cond[:, None, :].repeat(1, encoder_output.shape[1], 1)
# conditioning: encoder outputs + time embedding
condition = encoder_output + time_embed
condition = self.pos_drop(condition)
# target stream
trg_embed = self.trg_embed(trg_embed)
x = self.pe(trg_embed)
x = self.emb_dropout(x)
padding_mask = trg_mask
# causal mask for target self-attn
sub_mask = subsequent_mask(trg_embed.size(1)).type_as(trg_mask)
# cross-attend target stream to conditioning
x, _ = self.layers_pose_condition(
x=x,
memory=condition,
src_mask=src_mask,
trg_mask=sub_mask,
padding_mask=padding_mask,
)
# mid projection to split (joints vs bones) and re-embed
x = self.layer_norm_mid(x)
x = self.output_layer_mid(x) # [B,T,350]
o_reshaped = x.view(x.shape[0], x.shape[1], 50, 7)
o_1, o_2 = torch.split(o_reshaped, [3, 4], dim=-1) # joints(3) vs bones(4)
o_1 = o_1.reshape(o_1.shape[0], o_1.shape[1], 50 * 3)
o_2 = o_2.reshape(o_2.shape[0], o_2.shape[1], 50 * 4)
o_1 = self.o1_embed(o_1)
o_2 = self.o2_embed(o_2)
# second decoder layer mixes the two streams
x, _ = self.layers_mha_ac(
x=o_1,
memory=o_2,
src_mask=sub_mask,
trg_mask=sub_mask,
padding_mask=padding_mask,
)
# final norm + linear head -> joints x0
x = self.layer_norm(x)
output = self.output_layer(x) # [B,T,150]
return output
def __repr__(self):
return f"{self.__class__.__name__}(num_layers={self.num_layers}, num_heads={self.num_heads})"
|