VOSR / inference_tiles.py
41-807's picture
Upload inference_tiles.py with huggingface_hub
03233ed verified
Raw
History Blame Contribute Delete
8.48 kB
"""Latent-space tiled inference (from upstream VOSR scripts)."""
from __future__ import annotations
import torch
from pipeline import (
_decode_latent,
_encode_latent,
get_venc_features,
)
def _gaussian_weights(tile_h, tile_w, channels, device):
var = 0.01
mid_h, mid_w = (tile_h - 1) / 2, (tile_w - 1) / 2
y = torch.arange(tile_h, dtype=torch.float32)
x = torch.arange(tile_w, dtype=torch.float32)
wy = torch.exp(-((y - mid_h) / tile_h) ** 2 / (2 * var))
wx = torch.exp(-((x - mid_w) / tile_w) ** 2 / (2 * var))
w = wy[:, None] * wx[None, :]
return w.to(device).unsqueeze(0).unsqueeze(0).expand(1, channels, -1, -1)
def _make_tile_grid(length, tile, overlap):
stride = max(tile - overlap, 1)
if length <= tile:
return [0]
positions = list(range(0, length - tile + 1, stride))
if positions[-1] + tile < length:
positions.append(length - tile)
return sorted(set(positions))
def tiled_multistep(model, vosr_model, vae, venc, lq_tensor, args, device="cuda", light_decoder=None):
import logging
log = logging.getLogger("vosr")
AE_FACTOR = 8
b = lq_tensor.shape[0]
patch_size = getattr(args, "patch_size", 2)
log.info(
"tiled_multistep | lq=%s dit_tile=%s vae_tile=%s",
tuple(lq_tensor.shape),
getattr(args, "tile_size", 0),
getattr(args, "vae_tile_size", 0),
)
with torch.no_grad():
lq_latent, latents_mean, latents_std = _encode_latent(vae, lq_tensor, args, device)
log.info("tiled_multistep latent=%s", tuple(lq_latent.shape))
_, lc, lh, lw = lq_latent.shape
lt_size = max((args.tile_size // AE_FACTOR // patch_size) * patch_size, patch_size)
lt_overlap = max(args.tile_overlap // AE_FACTOR, lt_size // 8)
lt_size = min(lt_size, min(lh, lw))
lt_overlap = min(lt_overlap, lt_size - 1)
if lh <= lt_size and lw <= lt_size:
with torch.no_grad():
z_fea = get_venc_features(venc, lq_tensor, args) if venc is not None else None
sr_latent = vosr_model.sample_multistep_fm(
model, lq_latent, n_steps=args.infer_steps, venc_fea=z_fea
)
return _decode_latent(vae, sr_latent, args, latents_mean, latents_std, light_decoder)
h_pos = _make_tile_grid(lh, lt_size, lt_overlap)
w_pos = _make_tile_grid(lw, lt_size, lt_overlap)
g_weight = _gaussian_weights(lt_size, lt_size, lc, device)
tile_venc = {}
if venc is not None:
with torch.no_grad():
for hi in h_pos:
for wi in w_pos:
ph_s, pw_s = hi * AE_FACTOR, wi * AE_FACTOR
ph_e = min((hi + lt_size) * AE_FACTOR, lq_tensor.shape[2])
pw_e = min((wi + lt_size) * AE_FACTOR, lq_tensor.shape[3])
lq_crop = lq_tensor[:, :, ph_s:ph_e, pw_s:pw_e]
tile_venc[(hi, wi)] = get_venc_features(venc, lq_crop, args)
weak_str = (args.weak_cond_strength_aelq_list[0] + args.weak_cond_strength_aelq_list[1]) / 2.0
lq_weak_full = vosr_model.interpolate(
lq_latent, torch.zeros_like(lq_latent), weak_str, vosr_model.interp_type
)
z = torch.randn_like(lq_latent)
n_steps = args.infer_steps
t_seq = torch.linspace(1.0, 0.0, n_steps + 1, device=device)
cfg_scale = vosr_model.cfg_scale
t_start, t_end = vosr_model.t_start, vosr_model.t_end
with torch.no_grad():
for step_i in range(n_steps):
t_cur, t_nxt = t_seq[step_i], t_seq[step_i + 1]
dt = t_cur - t_nxt
use_cfg = t_cur >= t_start and t_cur <= t_end
u_acc = torch.zeros_like(lq_latent)
w_acc = torch.zeros_like(lq_latent)
for hi in h_pos:
for wi in w_pos:
he, we = hi + lt_size, wi + lt_size
lq_tile = lq_latent[:, :, hi:he, wi:we]
z_tile = z[:, :, hi:he, wi:we]
z_fea_tile = tile_venc.get((hi, wi))
inp_cond = torch.cat([lq_tile, z_tile], dim=1)
if use_cfg:
lq_weak_tile = lq_weak_full[:, :, hi:he, wi:we]
inp_weak = torch.cat([lq_weak_tile, z_tile], dim=1)
model_inp = torch.cat([inp_cond, inp_weak], dim=0)
model_t = t_cur.expand(b).repeat(2)
if z_fea_tile is not None:
model_z = [torch.cat([v, torch.zeros_like(v)], dim=0) for v in z_fea_tile]
else:
model_z = None
d_out = model(model_inp, model_t, z=model_z)
d_cond, d_weak = d_out.chunk(2)
u_tile = d_weak + cfg_scale * (d_cond - d_weak)
else:
u_tile = model(inp_cond, t_cur.expand(b), z=z_fea_tile)
u_acc[:, :, hi:he, wi:we] += u_tile * g_weight
w_acc[:, :, hi:he, wi:we] += g_weight
z = z - dt * (u_acc / w_acc)
with torch.no_grad():
return _decode_latent(vae, z, args, latents_mean, latents_std, light_decoder)
def tiled_onestep(model, vae, venc, lq_tensor, args, device="cuda", light_decoder=None):
import logging
log = logging.getLogger("vosr")
AE_FACTOR = 8
b = lq_tensor.shape[0]
patch_size = getattr(args, "patch_size", 2)
log.info(
"tiled_onestep | lq=%s dit_tile=%s vae_tile=%s",
tuple(lq_tensor.shape),
getattr(args, "tile_size", 0),
getattr(args, "vae_tile_size", 0),
)
with torch.no_grad():
lq_latent, latents_mean, latents_std = _encode_latent(vae, lq_tensor, args, device)
log.info("tiled_onestep latent=%s", tuple(lq_latent.shape))
_, lc, lh, lw = lq_latent.shape
lt_size = max((args.tile_size // AE_FACTOR // patch_size) * patch_size, patch_size)
lt_overlap = max(args.tile_overlap // AE_FACTOR, lt_size // 8)
lt_size = min(lt_size, min(lh, lw))
lt_overlap = min(lt_overlap, lt_size - 1)
if lh <= lt_size and lw <= lt_size:
with torch.no_grad():
z_fea = get_venc_features(venc, lq_tensor, args) if venc is not None else None
z = torch.randn_like(lq_latent)
n_steps = args.infer_steps
t_seq = torch.linspace(1.0, 0.0, n_steps + 1, device=device)
for i in range(n_steps):
t_cur, t_nxt = t_seq[i], t_seq[i + 1]
u = model(torch.cat([lq_latent, z], 1), t_cur.expand(b), t_nxt.expand(b), z_fea)
z = z - (t_cur - t_nxt) * u
return _decode_latent(vae, z, args, latents_mean, latents_std, light_decoder)
h_pos = _make_tile_grid(lh, lt_size, lt_overlap)
w_pos = _make_tile_grid(lw, lt_size, lt_overlap)
g_weight = _gaussian_weights(lt_size, lt_size, lc, device)
tile_venc = {}
if venc is not None:
with torch.no_grad():
for hi in h_pos:
for wi in w_pos:
ph_s, pw_s = hi * AE_FACTOR, wi * AE_FACTOR
ph_e = min((hi + lt_size) * AE_FACTOR, lq_tensor.shape[2])
pw_e = min((wi + lt_size) * AE_FACTOR, lq_tensor.shape[3])
lq_crop = lq_tensor[:, :, ph_s:ph_e, pw_s:pw_e]
tile_venc[(hi, wi)] = get_venc_features(venc, lq_crop, args)
z = torch.randn_like(lq_latent)
n_steps = args.infer_steps
t_seq = torch.linspace(1.0, 0.0, n_steps + 1, device=device)
with torch.no_grad():
for step_i in range(n_steps):
t_cur, t_nxt = t_seq[step_i], t_seq[step_i + 1]
dt = t_cur - t_nxt
u_acc = torch.zeros_like(lq_latent)
w_acc = torch.zeros_like(lq_latent)
for hi in h_pos:
for wi in w_pos:
he, we = hi + lt_size, wi + lt_size
inp = torch.cat(
[lq_latent[:, :, hi:he, wi:we], z[:, :, hi:he, wi:we]], dim=1
)
z_fea_tile = tile_venc.get((hi, wi))
u_tile = model(inp, t_cur.expand(b), t_nxt.expand(b), z_fea_tile)
u_acc[:, :, hi:he, wi:we] += u_tile * g_weight
w_acc[:, :, hi:he, wi:we] += g_weight
z = z - dt * (u_acc / w_acc)
with torch.no_grad():
return _decode_latent(vae, z, args, latents_mean, latents_std, light_decoder)