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