Diagonal Gaussian Dist
class DiagonalGaussianDistribution(object):
def __init__(self, parameters, deterministic=False, dim=1):
self.parameters = parameters
self.mean, self.logvar = torch.chunk(parameters, 2, dim=dim)
if dim == 1:
self.dims = [1, 2, 3]
elif dim == 2:
self.dims = [1, 2]
else:
raise NotImplementedError
self.logvar = torch.clamp(self.logvar, -30.0, 20.0)
self.deterministic = deterministic
self.std = torch.exp(0.5 * self.logvar)
self.var = torch.exp(self.logvar)
if self.deterministic:
self.var = self.std = torch.zeros_like(self.mean).to(device=self.parameters.device)
def sample(self):
x = self.mean + self.std * torch.randn(self.mean.shape).to(device=self.parameters.device)
return x
def mode(self):
return self.mean
sample point from normal_dist with given mean and std
Attend
rotary_freqs = RotaryEmbedding(
dim=head_dim // 4,
freqs_for="pixel",
max_freq=frame_height * frame_width,
).get_axial_freqs(frame_height, frame_width)
还是用了 rotary_emb for pixel
注意这里dim = head_dim // 4 ?
endow each pixel freq embed based on pos
Attend Block
self.norm1 = norm_layer(dim)
self.attn = Attention(
dim,
num_heads,
frame_height,
frame_width,
qkv_bias=qkv_bias,
)
self.norm2 = norm_layer(dim)
mlp_hidden_dim = int(dim * mlp_ratio)
self.mlp = Mlp(
in_features=dim,
hidden_features=mlp_hidden_dim,
act_layer=act_layer,
)
def forward(self, x):
x = x + self.attn(self.norm1(x))
x = x + self.mlp(self.norm2(x))
return x
注意这里不再用mod和activation了 只有norm
Autoencoder KL
input_image → patch_emb
encoder = enc_depth * AttentionBlock(enc_dim, enc_heads)
bottleneck:
self.quant_conv = nn.Linear(enc_dim, mult * latent_dim)
self.post_quant_conv = nn.Linear(latent_dim, dec_dim)
decoder = dec_depth * AttentionBlock(dec_dim, dec_heads)
self.predictor = nn.Linear(dec_dim, self.patch_dim)
*self.patch_dim = 3 * patch_size**2
encoder 完后
# bottleneck
moments = self.quant_conv(x) enc_dim -> mult * latent_dim
if not self.use_variational:
moments = torch.cat((moments, torch.zeros_like(moments)), 2)
posterior = DiagonalGaussianDistribution(moments, deterministic=(not self.use_variational), dim=2)
return posterior
从moments里sample 如果not variational → deterministic
var 设成跟moments.zeros_like 全部moments作为mean
若有var 把现有moments分成两份? 但是这样sample的x shape不会不同?
—> 因为如果没有var 不做sample
#autoencode
if self.use_variational and sample_posterior:
z = posterior.sample()
else:
z = posterior.mode() #only mean
decode 完后 接predictor
*self*.predictor = nn.Linear(dec_dim, *self*.patch_dim) *# decoder to patch
## self.patch_dim = 3 * patch_size**2*
然后unpatchify
def unpatchify(self, x):
bsz = x.shape[0]
# unpatchify
x = x.reshape(bsz, self.seq_h, self.seq_w, self.patch_dim).permute([0, 3, 1, 2])
## [b, h, w, cxpxp] --> [b, cxpxp, h, w]
x = x.reshape(
bsz,
3,
self.patch_size,
self.patch_size,
self.seq_h,
self.seq_w,
).permute([0, 1, 4, 2, 5, 3]) # [b, c, p, p, h, w] --> [b, c, h, p, w, p]
x = x.reshape(
bsz,
3,
self.input_height,
self.input_width,
) ## [b, c, hxp, wxp]
return x
最后返回type
def autoencode(self, input, sample_posterior=True):
posterior = self.encode(input)
if self.use_variational and sample_posterior:
z = posterior.sample()
else:
z = posterior.mode()
dec = self.decode(z)
return dec, posterior, z
def forward(self, inputs, labels, split="train"):
rec, post, latent = self.autoencode(inputs)
## returnable, posterior(encoder_result), latent(posterior sampled with norm)
return rec, post, latent
enc, dec 的dim, heads相同 dec depth两倍
360640 → patchify to size of 2020
def ViT_L_20_Shallow_Encoder(**kwargs):
if "latent_dim" in kwargs:
latent_dim = kwargs.pop("latent_dim")
else:
latent_dim = 16
return AutoencoderKL(
latent_dim=latent_dim,
patch_size=20,
enc_dim=1024,
enc_depth=6,
enc_heads=16,
dec_dim=1024,
dec_depth=12,
dec_heads=16,
input_height=360,
input_width=640,
**kwargs,
)
最近感想:
看着冰箱里剩的5块肉和4个辣椒陷入沉思
这波我只能说 拖就硬拖 但我就是要拖住!