"""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)