import torch from einops import rearrange EPS = 1e-6 class WarpMedian: def __init__(self, **kwargs): pass def warp(self, depth, **kwargs): if kwargs.get("reference_meta", None) is not None: median_val = kwargs["reference_meta"] return depth / torch.clamp(median_val, min=1e-3), (depth > EPS), median_val prompt_depth = kwargs.get("prompt_depth") prompt_mask = kwargs.get("prompt_mask") median_val = [] batch_size = depth.shape[0] for b in range(batch_size): valid_prompt = (prompt_mask[b] > 0.0) & torch.isfinite(prompt_depth[b]) & (prompt_depth[b] > 0.0) if not valid_prompt.any(): raise RuntimeError( "WarpMedian received an empty prompt set. Prompt-conditioned inputs " "must be filtered at the dataset stage before reaching InfiniDepth." ) median = torch.quantile(prompt_depth[b][valid_prompt], 0.5) median_val.append(median) median_val = torch.stack(median_val, dim=0) median_val = rearrange(median_val, "b -> b 1 1 1") return depth / torch.clamp(median_val, min=1e-2), (depth > EPS) & (prompt_mask > 0.0), median_val def unwarp(self, depth, **kwargs): median_val = kwargs.get("reference_meta") return depth * torch.clamp(median_val, min=1e-3)