import torch import torch.nn.functional as F import math from improved_tiling_functions import get_safe_epsilon # Caches _CUBEMAP_GRID_CACHE = {} _PANO_GRID_CACHE = {} _BLUR_KERNEL_CACHE = {} # ======================================================================== # CUBEMAP (3D) — Engine A (Fast) + Engine B (Seam-Blend) # ======================================================================== def _safe_pad4d(x, pad, mode='reflect', value=0.0): """ Safe wrapper around F.pad for 4D tensors. - For mode='reflect', PyTorch requires pad < input_size. If invalid, we fall back to 'replicate' to avoid runtime errors. pad: (left, right, top, bottom) """ if not isinstance(pad, (tuple, list)) or len(pad) != 4: return F.pad(x, pad, mode=mode, value=value) if mode == 'constant' else F.pad(x, pad, mode=mode) l, r, t, b = pad if mode == 'reflect': h = int(x.shape[-2]) w = int(x.shape[-1]) if (l >= w) or (r >= w) or (t >= h) or (b >= h): mode = 'replicate' if mode == 'constant': return F.pad(x, (l, r, t, b), mode=mode, value=value) return F.pad(x, (l, r, t, b), mode=mode) def _cubemap_split_faces(x): """ Splits a 3x2 cubemap net into faces. Layout expected (top row / bottom row): S | E | N B | T | W Returns tuple (S, E, N, B, T, W), each (B,C,h,w) """ B, C, H, W = x.shape if H % 2 != 0 or W % 3 != 0: raise ValueError("Cubemap expects H%2==0 and W%3==0 (3x2 net).") h, w = H // 2, W // 3 S = x[:, :, 0:h, 0:w] E = x[:, :, 0:h, w:2*w] N = x[:, :, 0:h, 2*w:3*w] Bm = x[:, :, h:2*h, 0:w] T = x[:, :, h:2*h, w:2*w] Wf = x[:, :, h:2*h, 2*w:3*w] return S, E, N, Bm, T, Wf def _cubemap_stitch_faces(S, E, N, Bm, T, Wf): """Stitches faces back into a 3x2 net (S/E/N over B/T/W).""" B, C, h, w = S.shape out = torch.zeros((B, C, h * 2, w * 3), device=S.device, dtype=S.dtype) out[:, :, 0:h, 0:w] = S out[:, :, 0:h, w:2*w] = E out[:, :, 0:h, 2*w:3*w] = N out[:, :, h:2*h, 0:w] = Bm out[:, :, h:2*h, w:2*w] = T out[:, :, h:2*h, 2*w:3*w] = Wf return out def _cubemap_pad_with_adjoint(O, L, R, U, D, pL, pR, pU, pD, pad_mode='replicate', seam_strength=0.0, seam_width=0): """ Pads a face O with neighbor strips L/R/U/D (already extracted from adjacent faces). Supports optional seam blending (Engine B) by mixing neighbor padding with O edge. """ B, C, h, w = O.shape Hp = h + pU + pD Wp = w + pL + pR Z = torch.zeros((B, C, Hp, Wp), device=O.device, dtype=O.dtype) Z[:, :, pU:pU + h, pL:pL + w] = O if pL == 0 and pR == 0 and pU == 0 and pD == 0: return Z # Helper: create ramp for seam_width (0 at boundary, 1 at outer pad) def _make_ramp(n, seam_w, device, dtype): if n <= 0: return None seam_w = int(max(0, min(seam_w, n))) if seam_w == 0: return torch.ones((n,), device=device, dtype=dtype) if seam_w == 1: ramp = torch.ones((n,), device=device, dtype=dtype) ramp[0] = 0.0 return ramp ramp = torch.ones((n,), device=device, dtype=dtype) ramp[:seam_w] = torch.linspace(0.0, 1.0, steps=seam_w, device=device, dtype=dtype) return ramp # Fill left/right strips if pL > 0: Lp = _safe_pad4d(L, (0, 0, pU, pD), mode=pad_mode) strip = Lp if seam_strength > 0.0: Oedge = O[:, :, :, :min(pL, w)] Oedge = _safe_pad4d(Oedge, (0, max(0, pL - Oedge.shape[-1]), pU, pD), mode='replicate') ramp = _make_ramp(pL, seam_width, O.device, O.dtype).view(1, 1, 1, pL) blend_scheme = Oedge * (1.0 - ramp) + strip * ramp strip = strip * (1.0 - seam_strength) + blend_scheme * seam_strength Z[:, :, :, :pL] = strip if pR > 0: Rp = _safe_pad4d(R, (0, 0, pU, pD), mode=pad_mode) strip = Rp if seam_strength > 0.0: Oedge = O[:, :, :, max(0, w - pR):w] need = pR - Oedge.shape[-1] Oedge = _safe_pad4d(Oedge, (max(0, need), 0, pU, pD), mode='replicate') ramp = _make_ramp(pR, seam_width, O.device, O.dtype).view(1, 1, 1, pR).flip(-1) blend_scheme = Oedge * (1.0 - ramp) + strip * ramp strip = strip * (1.0 - seam_strength) + blend_scheme * seam_strength Z[:, :, :, -pR:] = strip # Fill top/bottom strips if pU > 0: Up = _safe_pad4d(U, (pL, pR, 0, 0), mode=pad_mode) strip = Up if seam_strength > 0.0: Oedge = O[:, :, :min(pU, h), :] Oedge = _safe_pad4d(Oedge, (pL, pR, 0, max(0, pU - Oedge.shape[-2])), mode='replicate') ramp = _make_ramp(pU, seam_width, O.device, O.dtype).view(1, 1, pU, 1) blend_scheme = Oedge * (1.0 - ramp) + strip * ramp strip = strip * (1.0 - seam_strength) + blend_scheme * seam_strength Z[:, :, :pU, :] = strip if pD > 0: Dp = _safe_pad4d(D, (pL, pR, 0, 0), mode=pad_mode) strip = Dp if seam_strength > 0.0: Oedge = O[:, :, max(0, h - pD):h, :] need = pD - Oedge.shape[-2] Oedge = _safe_pad4d(Oedge, (pL, pR, max(0, need), 0), mode='replicate') ramp = _make_ramp(pD, seam_width, O.device, O.dtype).view(1, 1, pD, 1).flip(-2) blend_scheme = Oedge * (1.0 - ramp) + strip * ramp strip = strip * (1.0 - seam_strength) + blend_scheme * seam_strength Z[:, :, -pD:, :] = strip # Fix corners overlapping (same as cubemap(3).py logic) if pU and pL: Z[:, :, :pU, :pL] /= 2 if pU and pR: Z[:, :, :pU, -pR:] /= 2 if pD and pL: Z[:, :, -pD:, :pL] /= 2 if pD and pR: Z[:, :, -pD:, -pR:] /= 2 return Z def conv2d_cubemap_batched(input_tensor, weight, bias, stride, dilation, groups, pad_h, pad_w, pad_mode='replicate', engine='A (Fast)', seam_width=0, seam_strength=0.0): """ Cubemap convolution for a 3x2 cubemap net (S/E/N over B/T/W), using 1 conv call: - Engine A: neighbor padding (fast, like cubemap(3).py but batched) - Engine B: same, but with seam-aware blending inside padding regions. NOTE: Requires square faces (h == w) to keep rotations consistent. """ if pad_h != pad_w: return F.conv2d(input_tensor, weight, bias, stride, (pad_h, pad_w), dilation, groups) if pad_h == 0 and pad_w == 0: return F.conv2d(input_tensor, weight, bias, stride, (0, 0), dilation, groups) try: S, E, N, Bm, T, Wf = _cubemap_split_faces(input_tensor) except Exception: return F.conv2d(input_tensor, weight, bias, stride, (pad_h, pad_w), dilation, groups) B, C, h, w = S.shape if h != w: return F.conv2d(input_tensor, weight, bias, stride, (pad_h, pad_w), dilation, groups) p = int(pad_h) pL = pR = pU = pD = p seam_strength = float(max(0.0, min(seam_strength, 1.0))) if (engine or '').startswith('B') else 0.0 seam_width = int(max(0, seam_width)) ZS = _cubemap_pad_with_adjoint( S, L=Wf[:, :, :, -pL:], R=E[:, :, :, :pR], U=T[:, :, -pU:, :], D=Bm[:, :, :pD, :], pL=pL, pR=pR, pU=pU, pD=pD, pad_mode=pad_mode, seam_strength=seam_strength, seam_width=seam_width ) ZE = _cubemap_pad_with_adjoint( E, L=S[:, :, :, -pL:], R=N[:, :, :, :pR], U=torch.rot90(T[:, :, :, -pU:], k=-1, dims=[2, 3]), D=torch.rot90(Bm[:, :, :, -pD:], k=+1, dims=[2, 3]), pL=pL, pR=pR, pU=pU, pD=pD, pad_mode=pad_mode, seam_strength=seam_strength, seam_width=seam_width ) ZN = _cubemap_pad_with_adjoint( N, L=E[:, :, :, -pL:], R=Wf[:, :, :, :pR], U=T[:, :, :pU, :].flip(-1), D=Bm[:, :, -pD:, :].flip(-1), pL=pL, pR=pR, pU=pU, pD=pD, pad_mode=pad_mode, seam_strength=seam_strength, seam_width=seam_width ) ZB = _cubemap_pad_with_adjoint( Bm, L=torch.rot90(Wf[:, :, -pL:, :], k=+1, dims=[2, 3]), R=torch.rot90(E[:, :, -pR:, :], k=-1, dims=[2, 3]), U=S[:, :, -pU:, :], D=N[:, :, -pD:, :].flip(-1), pL=pL, pR=pR, pU=pU, pD=pD, pad_mode=pad_mode, seam_strength=seam_strength, seam_width=seam_width ) ZT = _cubemap_pad_with_adjoint( T, L=torch.rot90(Wf[:, :, :pL, :], k=-1, dims=[2, 3]), R=torch.rot90(E[:, :, :pR, :], k=+1, dims=[2, 3]), U=N[:, :, :pU, :].flip(-1), D=S[:, :, :pD, :], pL=pL, pR=pR, pU=pU, pD=pD, pad_mode=pad_mode, seam_strength=seam_strength, seam_width=seam_width ) ZW = _cubemap_pad_with_adjoint( Wf, L=N[:, :, :, -pL:], R=S[:, :, :, :pR], U=torch.rot90(T[:, :, :, :pL], k=+1, dims=[2, 3]), D=torch.rot90(Bm[:, :, :, :pD], k=-1, dims=[2, 3]), pL=pL, pR=pR, pU=pU, pD=pD, pad_mode=pad_mode, seam_strength=seam_strength, seam_width=seam_width ) Z = torch.cat([ZS, ZE, ZN, ZB, ZT, ZW], dim=0) Y = F.conv2d(Z, weight, bias, stride, (0, 0), dilation, groups) YS, YE, YN, YB, YT, YW = Y.chunk(6, dim=0) return _cubemap_stitch_faces(YS, YE, YN, YB, YT, YW) # =================================== # CUBEMAP (3D) — Engine C (GridSample / True 3D mapping) # =================================== def _ypr_rotation_matrix(yaw_deg: float, pitch_deg: float, roll_deg: float, device, dtype): """ Builds a rotation matrix from yaw/pitch/roll angles (degrees). Convention: - yaw around +Y axis - pitch around +X axis - roll around +Z axis Applied as: R = Rz(roll) @ Rx(pitch) @ Ry(yaw) """ yaw = math.radians(float(yaw_deg)) pitch = math.radians(float(pitch_deg)) roll = math.radians(float(roll_deg)) cy, sy = math.cos(yaw), math.sin(yaw) cp, sp = math.cos(pitch), math.sin(pitch) cr, sr = math.cos(roll), math.sin(roll) # Ry (yaw) Ry = torch.tensor([[cy, 0.0, sy], [0.0, 1.0, 0.0], [-sy, 0.0, cy]], device=device, dtype=dtype) # Rx (pitch) Rx = torch.tensor([[1.0, 0.0, 0.0], [0.0, cp, -sp], [0.0, sp, cp]], device=device, dtype=dtype) # Rz (roll) Rz = torch.tensor([[cr, -sr, 0.0], [sr, cr, 0.0], [0.0, 0.0, 1.0]], device=device, dtype=dtype) return (Rz @ Rx @ Ry) def _cubemap_dirs_from_face_uv(face_id: int, u, v): """ Maps face-local (u,v) to 3D direction vectors BEFORE normalization. Faces in our atlas mapping: 0: Front (+Z) -> S 1: Right (+X) -> E 2: Back (-Z) -> N 3: Bottom (-Y) -> Bm 4: Top (+Y) -> T 5: Left (-X) -> Wf u, v are broadcastable tensors, typically shaped (Hp, Wp) or (1,1,Hp,Wp) """ if face_id == 0: # +Z (Front) x, y, z = u, -v, torch.ones_like(u) elif face_id == 1: # +X (Right) x, y, z = torch.ones_like(u), -v, -u elif face_id == 2: # -Z (Back) x, y, z = -u, -v, -torch.ones_like(u) elif face_id == 3: # -Y (Bottom) x, y, z = u, -torch.ones_like(u), -v elif face_id == 4: # +Y (Top) x, y, z = u, torch.ones_like(u), v elif face_id == 5: # -X (Left) x, y, z = -torch.ones_like(u), -v, u else: raise ValueError("Invalid face_id for cubemap.") return x, y, z def _cubemap_dir_to_atlas_grid(x, y, z, face_h: int, face_w: int, device, dtype): """ Converts 3D direction vectors to a single atlas (3x2 net) sampling grid in [-1,1]. Returns grid shaped (..., 2) with last dim [x_norm, y_norm]. """ x = x.to(torch.float32) y = y.to(torch.float32) z = z.to(torch.float32) eps_val = get_safe_epsilon(torch.float32) eps = torch.tensor(eps_val, device=device, dtype=torch.float32) # Normalize directions (avoid divide-by-zero) inv_len = torch.rsqrt(torch.clamp(x * x + y * y + z * z, min=eps_val)) x = x * inv_len y = y * inv_len z = z * inv_len ax = x.abs() ay = y.abs() az = z.abs() # Major axis selection is_x = (ax >= ay) & (ax >= az) is_y = (ay >= ax) & (ay >= az) is_z = ~(is_x | is_y) # Face index map: 0..5 face_idx = torch.empty_like(x, dtype=torch.int64) # Defaults (placeholders) u = torch.zeros_like(x) v = torch.zeros_like(x) # +X / -X mask = is_x & (x >= 0) face_idx[mask] = 1 u[mask] = -z[mask] / (ax[mask] + eps) v[mask] = -y[mask] / (ax[mask] + eps) mask = is_x & (x < 0) face_idx[mask] = 5 u[mask] = z[mask] / (ax[mask] + eps) v[mask] = -y[mask] / (ax[mask] + eps) # +Y / -Y mask = is_y & (y >= 0) face_idx[mask] = 4 u[mask] = x[mask] / (ay[mask] + eps) v[mask] = z[mask] / (ay[mask] + eps) mask = is_y & (y < 0) face_idx[mask] = 3 u[mask] = x[mask] / (ay[mask] + eps) v[mask] = -z[mask] / (ay[mask] + eps) # +Z / -Z mask = is_z & (z >= 0) face_idx[mask] = 0 u[mask] = x[mask] / (az[mask] + eps) v[mask] = -y[mask] / (az[mask] + eps) mask = is_z & (z < 0) face_idx[mask] = 2 u[mask] = -x[mask] / (az[mask] + eps) v[mask] = -y[mask] / (az[mask] + eps) # Atlas tile offsets (col,row) for each face_idx # 0:F -> (0,0), 1:R -> (1,0), 2:B -> (2,0), 3:Bo -> (0,1), 4:T -> (1,1), 5:L -> (2,1) col = torch.zeros_like(u) row = torch.zeros_like(v) col = torch.where(face_idx == 0, torch.tensor(0.0, device=device, dtype=dtype), col) row = torch.where(face_idx == 0, torch.tensor(0.0, device=device, dtype=dtype), row) col = torch.where(face_idx == 1, torch.tensor(1.0, device=device, dtype=dtype), col) row = torch.where(face_idx == 1, torch.tensor(0.0, device=device, dtype=dtype), row) col = torch.where(face_idx == 2, torch.tensor(2.0, device=device, dtype=dtype), col) row = torch.where(face_idx == 2, torch.tensor(0.0, device=device, dtype=dtype), row) col = torch.where(face_idx == 3, torch.tensor(0.0, device=device, dtype=dtype), col) row = torch.where(face_idx == 3, torch.tensor(1.0, device=device, dtype=dtype), row) col = torch.where(face_idx == 4, torch.tensor(1.0, device=device, dtype=dtype), col) row = torch.where(face_idx == 4, torch.tensor(1.0, device=device, dtype=dtype), row) col = torch.where(face_idx == 5, torch.tensor(2.0, device=device, dtype=dtype), col) row = torch.where(face_idx == 5, torch.tensor(1.0, device=device, dtype=dtype), row) # Convert (u,v) [-1,1] -> atlas pixel coords -> normalized coords [-1,1] H_atlas = int(face_h * 2) W_atlas = int(face_w * 3) # align_corners=True mapping uses (W-1)/(H-1) x_pix = col * face_w + (u + 1.0) * 0.5 * (face_w - 1) y_pix = row * face_h + (v + 1.0) * 0.5 * (face_h - 1) x_norm = (x_pix / max(W_atlas - 1, 1)) * 2.0 - 1.0 y_norm = (y_pix / max(H_atlas - 1, 1)) * 2.0 - 1.0 grid = torch.stack([x_norm, y_norm], dim=-1).to(dtype) return grid def _build_cubemap_engine_c_grids(face_h: int, face_w: int, pad: int, yaw: float, pitch: float, roll: float, coord_mode: str = "Cartesian (Face UV)", twist_deg: float = 0.0, polar_scale: float = 1.0, polar_power: float = 1.0, swirl_deg: float = 0.0, swirl_power: float = 1.0, device=None, dtype=None, antipode: bool = False, angle_quant: float = 0.5): """ Builds and caches per-face sampling grids (Engine C) for cubemap atlas. Grids map each pixel in a padded face to the correct location in the 3x2 atlas. """ if face_h <= 1 or face_w <= 1: return None # Quantize angles to stabilize caching q = float(angle_quant) q_milli = int(round(float(q) * 1000.0)) if q_milli <= 0: q_milli = 1 yaw_t = int(round(float(yaw) / q)) pitch_t = int(round(float(pitch) / q)) roll_t = int(round(float(roll) / q)) twist_t = int(round(float(twist_deg) / q)) swirl_t = int(round(float(swirl_deg) / q)) yaw_q = float(yaw_t) * q pitch_q = float(pitch_t) * q roll_q = float(roll_t) * q twist_q = float(twist_t) * q swirl_q = float(swirl_t) * q # Quantize continuous params a bit for caching polar_scale_q = round(float(polar_scale) * 100.0) / 100.0 polar_power_q = round(float(polar_power) * 100.0) / 100.0 swirl_power_q = round(float(swirl_power) * 100.0) / 100.0 dev_type = getattr(device, "type", None) dev_index = getattr(device, "index", None) key = ( str(dev_type) if dev_type is not None else str(device), int(dev_index) if dev_index is not None else -1, str(dtype), int(face_h), int(face_w), int(pad), str(coord_mode), int(yaw_t), int(pitch_t), int(roll_t), int(twist_t), int(round(float(polar_scale_q) * 100.0)), int(round(float(polar_power_q) * 100.0)), int(swirl_t), int(round(float(swirl_power_q) * 100.0)), bool(antipode), int(q_milli)) cached = _CUBEMAP_GRID_CACHE.get(key, None) if cached is not None: return cached p = int(max(0, pad)) Hp = int(face_h + 2 * p) Wp = int(face_w + 2 * p) # Face-local u,v coordinate system (padded) j = torch.arange(Wp, device=device, dtype=dtype) i = torch.arange(Hp, device=device, dtype=dtype) denom_w = float(max(face_w - 1, 1)) denom_h = float(max(face_h - 1, 1)) u = 2.0 * ((j - p) / denom_w) - 1.0 v = 2.0 * ((i - p) / denom_h) - 1.0 # Broadcast to (Hp,Wp) u2 = u.view(1, Wp).expand(Hp, Wp) v2 = v.view(Hp, 1).expand(Hp, Wp) # Advanced UV transform (twist / polar warp / swirl) if coord_mode is None: coord_mode = "Cartesian (Face UV)" cm = str(coord_mode) twist_rad = float(twist_q) * (math.pi / 180.0) swirl_rad = float(swirl_q) * (math.pi / 180.0) do_polar = cm.startswith("Polar") if abs(twist_rad) > 1e-9 or abs(swirl_rad) > 1e-9 or do_polar: eps_val = get_safe_epsilon(dtype) r = torch.sqrt(u2 * u2 + v2 * v2 + eps_val) r_clamped = torch.clamp(r, 0.0, 2.0) theta = torch.atan2(v2, u2) theta = theta + twist_rad if abs(swirl_rad) > 1e-9: sp = float(swirl_power_q) theta = theta + swirl_rad * torch.pow(r_clamped, sp) if do_polar: ps = float(polar_scale_q) pp = float(polar_power_q) r2 = torch.pow(torch.clamp(r_clamped * ps, min=0.0), pp) else: r2 = r u2 = r2 * torch.cos(theta) v2 = r2 * torch.sin(theta) R = _ypr_rotation_matrix(yaw_q, pitch_q, roll_q, device=device, dtype=dtype) grids = [] for face_id in range(6): x, y, z = _cubemap_dirs_from_face_uv(face_id, u2, v2) # Rotate directions dirs = torch.stack([x, y, z], dim=-1) dirs = torch.matmul(dirs, R.transpose(0, 1)) if antipode: dirs = -dirs grid = _cubemap_dir_to_atlas_grid( dirs[..., 0], dirs[..., 1], dirs[..., 2], face_h=face_h, face_w=face_w, device=device, dtype=dtype ) grids.append(grid) grids = torch.stack(grids, dim=0) # (6,Hp,Wp,2) _CUBEMAP_GRID_CACHE[key] = grids return grids def _grid_sample_geoaa(atlas, grid, samples: int = 1, radius_px: float = 0.0, mode: str = "bilinear", padding_mode: str = "border"): """ Optional geometric AA (multi-sampling) for Engine C. - samples: 1..4 - radius_px: pixel radius in atlas space (approx) """ samples = int(max(1, min(int(samples), 4))) radius_px = float(max(0.0, radius_px)) # sanitize grid_sample args if mode not in ("bilinear", "nearest"): mode = "bilinear" if padding_mode not in ("border", "reflection", "zeros"): padding_mode = "border" if samples == 1 or radius_px <= 0.0: return F.grid_sample(atlas, grid, mode=mode, padding_mode=padding_mode, align_corners=True) B, C, H, W = atlas.shape # normalize radius to grid space (align_corners=True => 1px == 2/(W-1)) dx = (radius_px * 2.0) / max(W - 1, 1) dy = (radius_px * 2.0) / max(H - 1, 1) offsets = [(0.0, 0.0)] if samples >= 2: offsets.append((dx, dy)) if samples >= 3: offsets.append((-dx, dy)) if samples >= 4: offsets.append((dx, -dy)) acc = None for ox, oy in offsets: g = grid.clone() g[..., 0] = (g[..., 0] + ox).clamp(-1.0, 1.0) g[..., 1] = (g[..., 1] + oy).clamp(-1.0, 1.0) y = F.grid_sample(atlas, g, mode=mode, padding_mode=padding_mode, align_corners=True) acc = y if acc is None else (acc + y) return acc / float(len(offsets)) def conv2d_cubemap_gridsample(input_tensor, weight, bias, stride, dilation, groups, pad_h, pad_w, yaw=0.0, pitch=0.0, roll=0.0, coord_mode="Cartesian (Face UV)", twist_deg=0.0, polar_scale=1.0, polar_power=1.0, swirl_deg=0.0, swirl_power=1.0, grid_interp="bilinear", grid_padding="border", cache_angle_quant=0.5, geoaa_samples=1, geoaa_radius_px=0.0, antipode_strength=0.0): """ Engine C: True 3D cubemap mapping using grid_sample. - Builds padded faces by sampling from the full 3x2 atlas via direction mapping. - Supports yaw/pitch/roll rotation of the sampling directions. - Optional geometric AA (multi-sampling) and Kohaku-inspired antipode mixing. """ if pad_h != pad_w: return F.conv2d(input_tensor, weight, bias, stride, (pad_h, pad_w), dilation, groups) p = int(pad_h) if p <= 0: return F.conv2d(input_tensor, weight, bias, stride, (0, 0), dilation, groups) B, C, H, W = input_tensor.shape if H % 2 != 0 or W % 3 != 0: return F.conv2d(input_tensor, weight, bias, stride, (pad_h, pad_w), dilation, groups) face_h = H // 2 face_w = W // 3 if face_h != face_w: return F.conv2d(input_tensor, weight, bias, stride, (pad_h, pad_w), dilation, groups) device = input_tensor.device # grid_sample expects float grid; use float32 for stability if input is fp16/bf16 grid_dtype = torch.float32 if input_tensor.dtype in (torch.float16, torch.bfloat16) else input_tensor.dtype grids = _build_cubemap_engine_c_grids(face_h, face_w, p, yaw, pitch, roll, coord_mode, twist_deg, polar_scale, polar_power, swirl_deg, swirl_power, device, grid_dtype, antipode=False, angle_quant=cache_angle_quant) if grids is None: return F.conv2d(input_tensor, weight, bias, stride, (pad_h, pad_w), dilation, groups) antipode_strength = float(max(0.0, min(float(antipode_strength), 1.0))) if antipode_strength > 0.0: grids_anti = _build_cubemap_engine_c_grids(face_h, face_w, p, yaw, pitch, roll, coord_mode, twist_deg, polar_scale, polar_power, swirl_deg, swirl_power, device, grid_dtype, antipode=True, angle_quant=cache_angle_quant) else: grids_anti = None Hp = int(face_h + 2 * p) Wp = int(face_w + 2 * p) faces_padded = [] for face_id in range(6): g = grids[face_id].to(device=device) gB = g.unsqueeze(0).expand(B, Hp, Wp, 2).contiguous() y0 = _grid_sample_geoaa(input_tensor, gB, samples=geoaa_samples, radius_px=geoaa_radius_px, mode=grid_interp, padding_mode=grid_padding) if grids_anti is not None: ga = grids_anti[face_id].to(device=device) gaB = ga.unsqueeze(0).expand(B, Hp, Wp, 2).contiguous() y1 = _grid_sample_geoaa(input_tensor, gaB, samples=geoaa_samples, radius_px=geoaa_radius_px, mode=grid_interp, padding_mode=grid_padding) y0 = y0 * (1.0 - antipode_strength) + y1 * antipode_strength faces_padded.append(y0) Z = torch.cat(faces_padded, dim=0) # (6B,C,Hp,Wp) Y = F.conv2d(Z, weight, bias, stride, (0, 0), dilation, groups) YS, YE, YN, YB, YT, YW = Y.chunk(6, dim=0) return _cubemap_stitch_faces(YS, YE, YN, YB, YT, YW) # ======================================================================== # PANORAMA LIVE (Equirectangular) — Engine C (3D grid_sample) # ======================================================================== def _get_blur_kernel_1d(radius: int, device, dtype): """Depthwise 1D blur kernel along X (width).""" r = int(max(0, radius)) if r <= 0: return None k = 2 * r + 1 dev_type = getattr(device, "type", None) dev_index = getattr(device, "index", None) key = (int(k), str(dev_type) if dev_type is not None else str(device), int(dev_index) if dev_index is not None else -1, str(dtype)) ker = _BLUR_KERNEL_CACHE.get(key, None) if ker is not None: return ker w = torch.ones((k,), device=device, dtype=dtype) / float(k) ker = w.view(1, 1, 1, k) # (1,1,1,k) _BLUR_KERNEL_CACHE[key] = ker return ker def _apply_pole_blur_smoothing(x, strength: float = 0.0, radius: int = 0, power: float = 1.0): """ Applies circular horizontal blur near poles (top/bottom) with a smooth mask. x: (B,C,H,W) """ strength = float(max(0.0, min(float(strength), 1.0))) radius = int(max(0, int(radius))) power = float(max(0.25, min(float(power), 4.0))) if strength <= 0.0 or radius <= 0: return x B, C, H, W = x.shape device = x.device dtype = x.dtype ker = _get_blur_kernel_1d(radius, device, dtype) if ker is None: return x # Pole mask: 1 near top/bottom, 0 near equator yy = torch.linspace(0.0, 1.0, steps=H, device=device, dtype=dtype).view(1, 1, H, 1) t = torch.abs(yy - 0.5) * 2.0 # 0 at equator, 1 at poles pole_mask = torch.pow(torch.clamp(t, 0.0, 1.0), power) # (1,1,H,1) # Circular pad along X then depthwise conv xp = F.pad(x, (radius, radius, 0, 0), mode="circular") # Depthwise conv: expand kernel per-channel weight = ker.expand(C, 1, 1, ker.shape[-1]).contiguous() blurred = F.conv2d(xp, weight, bias=None, stride=1, padding=0, groups=C) m = pole_mask * strength return x * (1.0 - m) + blurred * m def _build_panorama_engine_c_grid(H: int, W: int, pad_h: int, pad_w: int, yaw: float, pitch: float, roll: float, coord_mode: str = "Cartesian (lon/lat)", polar_scale: float = 1.0, polar_power: float = 1.0, twist_deg: float = 0.0, twist_power: float = 1.0, swirl_deg: float = 0.0, swirl_power: float = 1.0, pole_ease_power: float = 1.0, antipode: bool = False, angle_quant: float = 0.5, device=None, dtype=None): """ Builds/caches a sampling grid for equirectangular panoramas. Grid maps output pixels in a padded canvas to source coords in the original panorama. Uses true 3D spherical mapping (yaw/pitch/roll) and optional UV warps. """ if H <= 1 or W <= 1: return None ph = int(max(0, pad_h)) pw = int(max(0, pad_w)) Hp = int(H + 2 * ph) Wp = int(W + 2 * pw) q = float(max(0.1, float(angle_quant))) q_milli = int(round(float(q) * 1000.0)) if q_milli <= 0: q_milli = 1 yaw_t = int(round(float(yaw) / q)) pitch_t = int(round(float(pitch) / q)) roll_t = int(round(float(roll) / q)) twist_t = int(round(float(twist_deg) / q)) swirl_t = int(round(float(swirl_deg) / q)) yaw_q = float(yaw_t) * q pitch_q = float(pitch_t) * q roll_q = float(roll_t) * q twist_q = float(twist_t) * q swirl_q = float(swirl_t) * q polar_scale_q = round(float(polar_scale) * 100.0) / 100.0 polar_power_q = round(float(polar_power) * 100.0) / 100.0 twist_power_q = round(float(twist_power) * 100.0) / 100.0 swirl_power_q = round(float(swirl_power) * 100.0) / 100.0 pole_ease_q = round(float(pole_ease_power) * 100.0) / 100.0 dev_type = getattr(device, "type", None) dev_index = getattr(device, "index", None) key = ( str(dev_type) if dev_type is not None else str(device), int(dev_index) if dev_index is not None else -1, str(dtype), int(H), int(W), int(ph), int(pw), str(coord_mode), int(yaw_t), int(pitch_t), int(roll_t), int(twist_t), int(round(float(twist_power_q) * 100.0)), int(swirl_t), int(round(float(swirl_power_q) * 100.0)), int(round(float(polar_scale_q) * 100.0)), int(round(float(polar_power_q) * 100.0)), int(round(float(pole_ease_q) * 100.0)), bool(antipode), int(q_milli)) cached = _PANO_GRID_CACHE.get(key, None) if cached is not None: return cached # Output pixel -> base lon/lat (can extend beyond [0,1] in padding; that's OK) j = torch.arange(Wp, device=device, dtype=dtype) i = torch.arange(Hp, device=device, dtype=dtype) denom_w = float(max(W - 1, 1)) denom_h = float(max(H - 1, 1)) u = (j - pw) / denom_w # 0..1 over original image v = (i - ph) / denom_h u2 = u.view(1, Wp).expand(Hp, Wp) v2 = v.view(Hp, 1).expand(Hp, Wp) # lon in radians (wrap naturally via sin/cos); lat in radians (can go beyond poles) lon = (u2 - 0.5) * (2.0 * math.pi) lat = (0.5 - v2) * math.pi cm = str(coord_mode or "Cartesian (lon/lat)") do_polar = cm.startswith("Polar") # --- Optional twist & swirl in (lon,lat) domain --- tr = float(twist_q) * (math.pi / 180.0) tp = float(max(0.25, min(float(twist_power_q), 4.0))) if abs(tr) > 1e-9: t = torch.clamp(torch.abs(lat) / (0.5 * math.pi), 0.0, 1.0) lon = lon + tr * torch.sign(lat) * torch.pow(t, tp) sr = float(swirl_q) * (math.pi / 180.0) sp = float(max(0.25, min(float(swirl_power_q), 4.0))) if abs(sr) > 1e-9: t = torch.clamp(torch.abs(lat) / (0.5 * math.pi), 0.0, 1.0) lon = lon + sr * torch.pow(t, sp) # --- Polar mode: radial warp around poles via latitude reparameterization --- if do_polar: ps = float(max(0.01, float(polar_scale_q))) pp = float(max(0.25, min(float(polar_power_q), 6.0))) # t=0 at equator, t=1 at poles t = torch.clamp(torch.abs(lat) / (0.5 * math.pi), 0.0, 1.0) r = 1.0 - t # r=1 at equator, 0 at poles r2 = torch.pow(torch.clamp(r * ps, min=0.0, max=1.0), pp) t2 = 1.0 - r2 lat = torch.sign(lat) * t2 * (0.5 * math.pi) # Convert (lon,lat) to 3D direction cl = torch.cos(lon) sl = torch.sin(lon) ca = torch.cos(lat) sa = torch.sin(lat) x = sl * ca y = sa z = cl * ca # Apply global rotation R = _ypr_rotation_matrix(yaw_q, pitch_q, roll_q, device=device, dtype=dtype) dirs = torch.stack([x, y, z], dim=-1) dirs = torch.matmul(dirs, R.transpose(0, 1)) if antipode: dirs = -dirs # Back to lon/lat x2 = dirs[..., 0] y2 = torch.clamp(dirs[..., 1], -1.0, 1.0) z2 = dirs[..., 2] lon2 = torch.atan2(x2, z2) # [-pi,pi] lat2 = torch.asin(y2) # [-pi/2,pi/2] # Pole easing curve (power) on latitude magnitude pe = float(max(0.25, min(float(pole_ease_q), 6.0))) if abs(pe - 1.0) > get_safe_epsilon(torch.float16): t = torch.clamp(torch.abs(lat2) / (0.5 * math.pi), 0.0, 1.0) t = torch.pow(t, pe) lat2 = torch.sign(lat2) * t * (0.5 * math.pi) # Convert to source UV [0,1) with X wrap u_src = (lon2 / (2.0 * math.pi)) + 0.5 u_src = torch.remainder(u_src, 1.0) # wrap horizontally v_src = 0.5 - (lat2 / math.pi) # 0..1 # to normalized grid_sample coords [-1,1] x_norm = u_src * 2.0 - 1.0 y_norm = v_src * 2.0 - 1.0 grid = torch.stack([x_norm, y_norm], dim=-1).to(dtype) # (Hp,Wp,2) _PANO_GRID_CACHE[key] = grid return grid def conv2d_panorama_gridsample(input_tensor, weight, bias, stride, dilation, groups, pad_h, pad_w, yaw=0.0, pitch=0.0, roll=0.0, coord_mode="Cartesian (lon/lat)", polar_scale=1.0, polar_power=1.0, twist_deg=0.0, twist_power=1.0, swirl_deg=0.0, swirl_power=1.0, pole_ease_power=1.0, grid_interp="bilinear", grid_padding="border", cache_angle_quant=0.5, geoaa_samples=1, geoaa_radius_px=0.0, antipode_strength=0.0, pole_blur_strength=0.0, pole_blur_radius=0, pole_blur_power=1.0): """ Panorama Live Engine C: - Builds a padded panorama by sampling the original via 3D spherical mapping. - Runs conv2d without extra padding. - Optional Kohaku-style antipode mixing and pole blur smoothing. """ ph = int(max(0, int(pad_h))) pw = int(max(0, int(pad_w))) if ph <= 0 and pw <= 0: return F.conv2d(input_tensor, weight, bias, stride, (0, 0), dilation, groups) B, C, H, W = input_tensor.shape device = input_tensor.device grid_dtype = torch.float32 if input_tensor.dtype in (torch.float16, torch.bfloat16) else input_tensor.dtype grid = _build_panorama_engine_c_grid( H, W, ph, pw, yaw=yaw, pitch=pitch, roll=roll, coord_mode=coord_mode, polar_scale=polar_scale, polar_power=polar_power, twist_deg=twist_deg, twist_power=twist_power, swirl_deg=swirl_deg, swirl_power=swirl_power, pole_ease_power=pole_ease_power, antipode=False, angle_quant=cache_angle_quant, device=device, dtype=grid_dtype ) if grid is None: return F.conv2d(input_tensor, weight, bias, stride, (pad_h, pad_w), dilation, groups) Hp = int(H + 2 * ph) Wp = int(W + 2 * pw) gB = grid.unsqueeze(0).expand(B, Hp, Wp, 2).contiguous() y0 = _grid_sample_geoaa(input_tensor, gB, samples=geoaa_samples, radius_px=geoaa_radius_px, mode=grid_interp, padding_mode=grid_padding) antipode_strength = float(max(0.0, min(float(antipode_strength), 1.0))) if antipode_strength > 0.0: grid_a = _build_panorama_engine_c_grid( H, W, ph, pw, yaw=yaw, pitch=pitch, roll=roll, coord_mode=coord_mode, polar_scale=polar_scale, polar_power=polar_power, twist_deg=twist_deg, twist_power=twist_power, swirl_deg=swirl_deg, swirl_power=swirl_power, pole_ease_power=pole_ease_power, antipode=True, angle_quant=cache_angle_quant, device=device, dtype=grid_dtype ) gaB = grid_a.unsqueeze(0).expand(B, Hp, Wp, 2).contiguous() y1 = _grid_sample_geoaa(input_tensor, gaB, samples=geoaa_samples, radius_px=geoaa_radius_px, mode=grid_interp, padding_mode=grid_padding) y0 = y0 * (1.0 - antipode_strength) + y1 * antipode_strength # Optional pole blur y0 = _apply_pole_blur_smoothing(y0, strength=pole_blur_strength, radius=pole_blur_radius, power=pole_blur_power) return F.conv2d(y0, weight, bias, stride, (0, 0), dilation, groups)