gestalt / ibq_decoder.py
multimodalart's picture
multimodalart HF Staff
Upload folder using huggingface_hub
69a2568 verified
Raw History Blame Contribute Delete
10.8 kB
"""Minimal IBQ-Tokenizer decoder for Gestalt text-to-image generation.
Vendors only the decode path of the IBQ vision tokenizer
(TencentARC/IBQ-Tokenizer-16384, from the SEED-Voken codebase, Apache-2.0):
the VQGAN-style ``Decoder`` architecture plus the codebook embedding and
``post_quant_conv``. Weights are loaded from the published Lightning
checkpoint, preferring the EMA parameters (the authors decode under
``ema_scope()``).
Decode semantics match SEED-Voken exactly:
indices (b, h*w)
-> get_codebook_entry(indices, shape=(b, h, w, c))
-> post_quant_conv -> decoder -> clamp(-1, 1) -> PIL image.
"""
import numpy as np
import torch
import torch.nn as nn
def nonlinearity(x):
# swish
return x * torch.sigmoid(x)
def Normalize(in_channels):
return torch.nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True)
class Upsample(nn.Module):
def __init__(self, in_channels, with_conv):
super().__init__()
self.with_conv = with_conv
if self.with_conv:
self.conv = torch.nn.Conv2d(in_channels, in_channels, kernel_size=3, stride=1, padding=1)
def forward(self, x):
x = torch.nn.functional.interpolate(x, scale_factor=2.0, mode="nearest")
if self.with_conv:
x = self.conv(x)
return x
class ResnetBlock(nn.Module):
def __init__(self, *, in_channels, out_channels=None, conv_shortcut=False,
dropout, temb_channels=512):
super().__init__()
self.in_channels = in_channels
out_channels = in_channels if out_channels is None else out_channels
self.out_channels = out_channels
self.use_conv_shortcut = conv_shortcut
self.norm1 = Normalize(in_channels)
self.conv1 = torch.nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1)
if temb_channels > 0:
self.temb_proj = torch.nn.Linear(temb_channels, out_channels)
self.norm2 = Normalize(out_channels)
self.dropout = torch.nn.Dropout(dropout)
self.conv2 = torch.nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1)
if self.in_channels != self.out_channels:
if self.use_conv_shortcut:
self.conv_shortcut = torch.nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1)
else:
self.nin_shortcut = torch.nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=1, padding=0)
def forward(self, x, temb):
h = x
h = self.norm1(h)
h = nonlinearity(h)
h = self.conv1(h)
if temb is not None:
h = h + self.temb_proj(nonlinearity(temb))[:, :, None, None]
h = self.norm2(h)
h = nonlinearity(h)
h = self.dropout(h)
h = self.conv2(h)
if self.in_channels != self.out_channels:
if self.use_conv_shortcut:
x = self.conv_shortcut(x)
else:
x = self.nin_shortcut(x)
return x + h
class AttnBlock(nn.Module):
def __init__(self, in_channels):
super().__init__()
self.in_channels = in_channels
self.norm = Normalize(in_channels)
self.q = torch.nn.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0)
self.k = torch.nn.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0)
self.v = torch.nn.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0)
self.proj_out = torch.nn.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0)
def forward(self, x):
h_ = x
h_ = self.norm(h_)
q = self.q(h_)
k = self.k(h_)
v = self.v(h_)
# compute attention
b, c, h, w = q.shape
q = q.reshape(b, c, h * w)
q = q.permute(0, 2, 1) # b, hw, c
k = k.reshape(b, c, h * w) # b, c, hw
w_ = torch.bmm(q, k) # b, hw, hw w[b,i,j]=sum_c q[b,i,c]k[b,c,j]
w_ = w_ * (int(c) ** (-0.5))
w_ = torch.nn.functional.softmax(w_, dim=2)
# attend to values
v = v.reshape(b, c, h * w)
w_ = w_.permute(0, 2, 1) # b, hw, hw (first hw of k, second of q)
h_ = torch.bmm(v, w_) # b, c, hw (hw of q) h_[b,c,j] = sum_i v[b,c,i] w_[b,i,j]
h_ = h_.reshape(b, c, h, w)
h_ = self.proj_out(h_)
return x + h_
class Decoder(nn.Module):
def __init__(self, *, ch, out_ch, ch_mult=(1, 2, 4, 8), num_res_blocks,
attn_resolutions, dropout=0.0, resamp_with_conv=True, in_channels,
resolution, z_channels, give_pre_end=False, **ignorekwargs):
super().__init__()
self.ch = ch
self.temb_ch = 0
self.num_resolutions = len(ch_mult)
self.num_res_blocks = num_res_blocks
self.resolution = resolution
self.in_channels = in_channels
self.give_pre_end = give_pre_end
# compute in_ch_mult, block_in and curr_res at lowest res
in_ch_mult = (1,) + tuple(ch_mult)
block_in = ch * ch_mult[self.num_resolutions - 1]
curr_res = resolution // 2 ** (self.num_resolutions - 1)
self.z_shape = (1, z_channels, curr_res, curr_res)
# z to block_in
self.conv_in = torch.nn.Conv2d(z_channels, block_in, kernel_size=3, stride=1, padding=1)
# middle
self.mid = nn.Module()
self.mid.block_1 = ResnetBlock(in_channels=block_in, out_channels=block_in,
temb_channels=self.temb_ch, dropout=dropout)
self.mid.attn_1 = AttnBlock(block_in)
self.mid.block_2 = ResnetBlock(in_channels=block_in, out_channels=block_in,
temb_channels=self.temb_ch, dropout=dropout)
# upsampling
self.up = nn.ModuleList()
for i_level in reversed(range(self.num_resolutions)):
block = nn.ModuleList()
attn = nn.ModuleList()
block_out = ch * ch_mult[i_level]
for i_block in range(self.num_res_blocks + 1):
block.append(ResnetBlock(in_channels=block_in, out_channels=block_out,
temb_channels=self.temb_ch, dropout=dropout))
block_in = block_out
if curr_res in attn_resolutions:
attn.append(AttnBlock(block_in))
up = nn.Module()
up.block = block
up.attn = attn
if i_level != 0:
up.upsample = Upsample(block_in, resamp_with_conv)
curr_res = curr_res * 2
self.up.insert(0, up) # prepend to get consistent order
# end
self.norm_out = Normalize(block_in)
self.conv_out = torch.nn.Conv2d(block_in, out_ch, kernel_size=3, stride=1, padding=1)
def forward(self, z):
# timestep embedding (unused, unconditional decoder)
temb = None
# z to block_in
h = self.conv_in(z)
# middle
h = self.mid.block_1(h, temb)
h = self.mid.attn_1(h)
h = self.mid.block_2(h, temb)
# upsampling
for i_level in reversed(range(self.num_resolutions)):
for i_block in range(self.num_res_blocks + 1):
h = self.up[i_level].block[i_block](h, temb)
if len(self.up[i_level].attn) > 0:
h = self.up[i_level].attn[i_block](h)
if i_level != 0:
h = self.up[i_level].upsample(h)
# end
h = self.norm_out(h)
h = nonlinearity(h)
h = self.conv_out(h)
return h
class IBQDecoder(nn.Module):
"""Codebook embedding + post_quant_conv + VQGAN decoder for IBQ tokens.
Matches ``TencentARC/IBQ-Tokenizer-16384`` (imagenet256_16384.ckpt,
n_embed=16384, embed_dim=256, z_channels=256, 256x256 native resolution).
"""
DDCONFIG = dict(
double_z=False,
z_channels=256,
resolution=256,
in_channels=3,
out_ch=3,
ch=128,
ch_mult=[1, 1, 2, 2, 4],
num_res_blocks=4,
attn_resolutions=[16],
dropout=0.0,
)
def __init__(self):
super().__init__()
cfg = self.DDCONFIG
self.decoder = Decoder(**cfg)
self.quantize_embedding = nn.Embedding(16384, 256)
self.post_quant_conv = nn.Conv2d(256, cfg["z_channels"], 1)
def load_from_checkpoint(self, ckpt_path: str):
"""Load EMA weights (fallback: plain) from the Lightning .ckpt."""
sd = torch.load(ckpt_path, map_location="cpu", weights_only=False)["state_dict"]
target_keys = set(self.state_dict().keys())
new_params = {}
ema_key_map = {}
# First pass: build EMA-name map ("model_ema.<dots-removed>" -> plain name)
for k in sd.keys():
s_name = k.replace(".", "")
if k.startswith("model_ema."):
ema_key_map[s_name] = k
# Second pass: prefer EMA values, fall back to plain values
for k in target_keys:
if k.startswith("quantize_embedding."):
base = "quantize.embedding." + k.split(".", 1)[1]
else:
base = k
ema_name = "model_ema." + base.replace(".", "")
if ema_name in sd:
new_params[k] = sd[ema_name]
elif base in sd:
new_params[k] = sd[base]
missing = target_keys - set(new_params.keys())
if missing:
raise ValueError(f"IBQ checkpoint missing keys: {sorted(missing)[:5]}")
self.load_state_dict(new_params, strict=True)
return self
def decode_indices(self, indices: torch.Tensor, lat_h: int = 32, lat_w: int = 32) -> torch.Tensor:
"""IBQ codebook indices (b, h*w) -> decoded image tensor (b, 3, H, W) in [-1, 1].
Mirrors ``IndexPropagationQuantize.get_codebook_entry`` +
``VQModel.decode`` from SEED-Voken.
"""
b = indices.shape[0]
c = self.quantize_embedding.embedding_dim
z_q = self.quantize_embedding(indices) # (b, h*w, c)
z_q = z_q.view(b, lat_h, lat_w, c) # shape=(b, h, w, c)
z_q = z_q.permute(0, 3, 1, 2).contiguous() # (b, c, h, w)
quant = self.post_quant_conv(z_q)
return self.decoder(quant)
@staticmethod
def to_pil(x: torch.Tensor) -> "Image.Image":
"""SEED-Voken ``custom_to_pil``: [-1,1] CHW tensor -> PIL RGB image."""
from PIL import Image
x = x.detach().cpu()
x = torch.clamp(x, -1.0, 1.0)
x = (x + 1.0) / 2.0
x = x.permute(1, 2, 0).numpy()
x = (255 * x).astype(np.uint8)
return Image.fromarray(x)