Spaces:
Running on Zero
Running on Zero
| 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) | |