Add HI-Mapper source (code only, no weights)
Browse files- .gitattributes +2 -32
- README.md +92 -0
- example.py +10 -0
- hi_mapper/ARCHITECTURE.md +340 -0
- hi_mapper/__init__.py +26 -0
- hi_mapper/hi_mapper.py +198 -0
- hi_mapper/hyp_diffusion.py +52 -0
- hi_mapper/lorentz.py +320 -0
- hi_mapper/tree.py +165 -0
- requirements.txt +1 -0
.gitattributes
CHANGED
|
@@ -1,35 +1,5 @@
|
|
| 1 |
-
*.7z filter=lfs diff=lfs merge=lfs -text
|
| 2 |
-
*.arrow filter=lfs diff=lfs merge=lfs -text
|
| 3 |
-
*.bin filter=lfs diff=lfs merge=lfs -text
|
| 4 |
-
*.bz2 filter=lfs diff=lfs merge=lfs -text
|
| 5 |
-
*.ckpt filter=lfs diff=lfs merge=lfs -text
|
| 6 |
-
*.ftz filter=lfs diff=lfs merge=lfs -text
|
| 7 |
-
*.gz filter=lfs diff=lfs merge=lfs -text
|
| 8 |
-
*.h5 filter=lfs diff=lfs merge=lfs -text
|
| 9 |
-
*.joblib filter=lfs diff=lfs merge=lfs -text
|
| 10 |
-
*.lfs.* filter=lfs diff=lfs merge=lfs -text
|
| 11 |
-
*.mlmodel filter=lfs diff=lfs merge=lfs -text
|
| 12 |
-
*.model filter=lfs diff=lfs merge=lfs -text
|
| 13 |
-
*.msgpack filter=lfs diff=lfs merge=lfs -text
|
| 14 |
-
*.npy filter=lfs diff=lfs merge=lfs -text
|
| 15 |
-
*.npz filter=lfs diff=lfs merge=lfs -text
|
| 16 |
-
*.onnx filter=lfs diff=lfs merge=lfs -text
|
| 17 |
-
*.ot filter=lfs diff=lfs merge=lfs -text
|
| 18 |
-
*.parquet filter=lfs diff=lfs merge=lfs -text
|
| 19 |
-
*.pb filter=lfs diff=lfs merge=lfs -text
|
| 20 |
-
*.pickle filter=lfs diff=lfs merge=lfs -text
|
| 21 |
-
*.pkl filter=lfs diff=lfs merge=lfs -text
|
| 22 |
*.pt filter=lfs diff=lfs merge=lfs -text
|
| 23 |
*.pth filter=lfs diff=lfs merge=lfs -text
|
| 24 |
-
*.
|
| 25 |
*.safetensors filter=lfs diff=lfs merge=lfs -text
|
| 26 |
-
|
| 27 |
-
*.tar.* filter=lfs diff=lfs merge=lfs -text
|
| 28 |
-
*.tar filter=lfs diff=lfs merge=lfs -text
|
| 29 |
-
*.tflite filter=lfs diff=lfs merge=lfs -text
|
| 30 |
-
*.tgz filter=lfs diff=lfs merge=lfs -text
|
| 31 |
-
*.wasm filter=lfs diff=lfs merge=lfs -text
|
| 32 |
-
*.xz filter=lfs diff=lfs merge=lfs -text
|
| 33 |
-
*.zip filter=lfs diff=lfs merge=lfs -text
|
| 34 |
-
*.zst filter=lfs diff=lfs merge=lfs -text
|
| 35 |
-
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
*.pt filter=lfs diff=lfs merge=lfs -text
|
| 2 |
*.pth filter=lfs diff=lfs merge=lfs -text
|
| 3 |
+
*.bin filter=lfs diff=lfs merge=lfs -text
|
| 4 |
*.safetensors filter=lfs diff=lfs merge=lfs -text
|
| 5 |
+
*.ckpt filter=lfs diff=lfs merge=lfs -text
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
README.md
ADDED
|
@@ -0,0 +1,92 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
---
|
| 2 |
+
license: apache-2.0
|
| 3 |
+
library_name: pytorch
|
| 4 |
+
tags:
|
| 5 |
+
- hyperbolic
|
| 6 |
+
- lorentz
|
| 7 |
+
- hierarchy
|
| 8 |
+
- entailment
|
| 9 |
+
- pedestrian-attribute-recognition
|
| 10 |
+
- clip
|
| 11 |
+
- promptpar
|
| 12 |
+
- hi-mapper
|
| 13 |
+
pipeline_tag: feature-extraction
|
| 14 |
+
---
|
| 15 |
+
|
| 16 |
+
# HI-Mapper (code only)
|
| 17 |
+
|
| 18 |
+
Hyperbolic hierarchy mapper for pedestrian attribute recognition (PromptPAR).
|
| 19 |
+
This Hub repo contains **source code only** — **no trained weights**.
|
| 20 |
+
|
| 21 |
+
HI-Mapper lifts Euclidean region / CLIP features into the Lorentz hyperboloid,
|
| 22 |
+
builds a fixed depth-3 anatomical tree, and supervises it with entailment cones,
|
| 23 |
+
sibling separation, and radius ordering (MERU / HyCoCLIP style). An optional
|
| 24 |
+
attribute-grounded entailment term maps PETA-style attribute prefixes onto tree
|
| 25 |
+
nodes. An optional HypDAE-style hyperbolic diffusion decoder is included.
|
| 26 |
+
|
| 27 |
+
## Contents
|
| 28 |
+
|
| 29 |
+
```
|
| 30 |
+
hi_mapper/
|
| 31 |
+
__init__.py # public exports
|
| 32 |
+
lorentz.py # Lorentz manifold + EuclideanToLorentz lift
|
| 33 |
+
tree.py # hierarchical / attribute entailment losses
|
| 34 |
+
hi_mapper.py # DivHiMapper + PETA attr grouping
|
| 35 |
+
hyp_diffusion.py # optional hyperbolic diffusion decoder
|
| 36 |
+
ARCHITECTURE.md # detailed architecture notes
|
| 37 |
+
```
|
| 38 |
+
|
| 39 |
+
## Install
|
| 40 |
+
|
| 41 |
+
```bash
|
| 42 |
+
pip install torch
|
| 43 |
+
# then copy hi_mapper/ into your project, or:
|
| 44 |
+
git clone https://huggingface.co/ZACK777/hi-mapper
|
| 45 |
+
```
|
| 46 |
+
|
| 47 |
+
Requires PyTorch. No Hub weights are downloaded.
|
| 48 |
+
|
| 49 |
+
## Quick start
|
| 50 |
+
|
| 51 |
+
```python
|
| 52 |
+
import torch
|
| 53 |
+
from hi_mapper import DivHiMapper, build_attr_groups
|
| 54 |
+
|
| 55 |
+
# region tokens: [B, 5, D] = global + 4 anatomical leaves
|
| 56 |
+
B, D = 2, 768
|
| 57 |
+
region_tokens = torch.randn(B, 5, D)
|
| 58 |
+
cls = torch.randn(B, D)
|
| 59 |
+
|
| 60 |
+
mapper = DivHiMapper(
|
| 61 |
+
feat_dim=D,
|
| 62 |
+
curvature=0.2,
|
| 63 |
+
target_radius=1.0,
|
| 64 |
+
attr_groups=None, # or build_attr_groups(attr_names)
|
| 65 |
+
)
|
| 66 |
+
root, mid, leaves, hier_loss, prompt_loss, attr_loss = mapper(
|
| 67 |
+
region_tokens, cls
|
| 68 |
+
)
|
| 69 |
+
```
|
| 70 |
+
|
| 71 |
+
## Geometry notes
|
| 72 |
+
|
| 73 |
+
- `EuclideanToLorentz` uses a running-norm scaler (not LayerNorm) so relative
|
| 74 |
+
norms — and thus hyperbolic radii — stay meaningful.
|
| 75 |
+
- Distance uses a numerically stable `asinh` form; Minkowski products run in
|
| 76 |
+
float64 intermediates.
|
| 77 |
+
- Entailment uses half-aperture / exterior-angle cones (`K=0.1`).
|
| 78 |
+
|
| 79 |
+
## Integration
|
| 80 |
+
|
| 81 |
+
Designed as a drop-in branch for PromptPAR / CLIP-based PAR models. See
|
| 82 |
+
`hi_mapper/ARCHITECTURE.md` for the tree layout, losses, and optional decoder.
|
| 83 |
+
|
| 84 |
+
## Citation / context
|
| 85 |
+
|
| 86 |
+
Built for hyperbolic hierarchical regularisation on top of PromptPAR-style
|
| 87 |
+
region tokens. Related ideas: MERU, HyCoCLIP, HypDAE, Khrulkov et al. (δ-hyperbolicity).
|
| 88 |
+
|
| 89 |
+
## License
|
| 90 |
+
|
| 91 |
+
Apache-2.0 (code). Upstream PromptPAR / CLIP licenses still apply if you use
|
| 92 |
+
those backbones.
|
example.py
ADDED
|
@@ -0,0 +1,10 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Minimal smoke test for HI-Mapper (no weights)."""
|
| 2 |
+
import torch
|
| 3 |
+
from hi_mapper import DivHiMapper
|
| 4 |
+
|
| 5 |
+
B, D = 2, 768
|
| 6 |
+
mapper = DivHiMapper(feat_dim=D, curvature=0.2, target_radius=1.0)
|
| 7 |
+
root, mid, leaves, hier_loss, prompt_loss, attr_loss = mapper(
|
| 8 |
+
torch.randn(B, 5, D), torch.randn(B, D)
|
| 9 |
+
)
|
| 10 |
+
print("root", tuple(root.shape), "hier_loss", float(hier_loss))
|
hi_mapper/ARCHITECTURE.md
ADDED
|
@@ -0,0 +1,340 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Hi-Mapper + Hyperbolic Generative Decoder
|
| 2 |
+
|
| 3 |
+
Detailed architecture of the **Hi-Mapper** plug-in module: hyperbolic hierarchy mapping and an optional HypDAE-style generative decoder.
|
| 4 |
+
|
| 5 |
+
---
|
| 6 |
+
|
| 7 |
+
## Overview
|
| 8 |
+
|
| 9 |
+
The module takes a set of **region-level visual tokens** and builds a **depth-3 binary tree** in Lorentz hyperbolic space. It produces:
|
| 10 |
+
|
| 11 |
+
1. A **refined root feature** for downstream use.
|
| 12 |
+
2. A **hierarchical contrastive loss** that enforces parent–child and sibling relations between levels.
|
| 13 |
+
3. Optionally, a **prompt-hierarchy alignment loss** when auxiliary prompt tokens are provided.
|
| 14 |
+
4. Optionally, a **diffusion reconstruction loss** that treats the hierarchy as a generative code for a global feature (train only).
|
| 15 |
+
|
| 16 |
+
```
|
| 17 |
+
region_tokens [B, 5, D] global token + 4 region leaves
|
| 18 |
+
global_feat [B, D] optional anchor (e.g. CLS)
|
| 19 |
+
prompt_nodes [vis_depth, …] optional, for prompt alignment
|
| 20 |
+
│
|
| 21 |
+
▼
|
| 22 |
+
┌─────────────────────────────┐
|
| 23 |
+
│ DivHiMapper │
|
| 24 |
+
│ Euclidean tree → Lorentz │
|
| 25 |
+
│ hierarchical contrastive │
|
| 26 |
+
│ optional prompt alignment │
|
| 27 |
+
└─────────────────────────────┘
|
| 28 |
+
│
|
| 29 |
+
├── root_feat [B, D]
|
| 30 |
+
├── mid_feats [B, 2, D]
|
| 31 |
+
├── leaf_feats [B, 4, D]
|
| 32 |
+
├── hier_loss
|
| 33 |
+
└── prompt_loss (optional)
|
| 34 |
+
│
|
| 35 |
+
▼
|
| 36 |
+
┌─────────────────────────────┐
|
| 37 |
+
│ HyperbolicDiffusionDecoder │ (train only, optional)
|
| 38 |
+
│ z = hierarchy code → CLS │
|
| 39 |
+
└─────────────────────────────┘
|
| 40 |
+
│
|
| 41 |
+
└── diff_loss
|
| 42 |
+
```
|
| 43 |
+
|
| 44 |
+
**Combined module loss:**
|
| 45 |
+
|
| 46 |
+
```
|
| 47 |
+
hi_mapper_loss = hier_loss
|
| 48 |
+
+ prompt_hier_w * prompt_loss # optional
|
| 49 |
+
+ hyp_diffusion_w * diff_loss # optional
|
| 50 |
+
```
|
| 51 |
+
|
| 52 |
+
---
|
| 53 |
+
|
| 54 |
+
## 1. Inputs and outputs
|
| 55 |
+
|
| 56 |
+
### DivHiMapper (`hi_mapper.py`)
|
| 57 |
+
|
| 58 |
+
| Tensor | Shape | Role |
|
| 59 |
+
|--------|-------|------|
|
| 60 |
+
| `region_tokens` | `[B, 5, D]` | Index 0 = global region token; indices 1–4 = four anatomical leaves (head, upper, lower, feet) |
|
| 61 |
+
| `global_feat` | `[B, D]` | Optional global anchor (soft-combined into root) |
|
| 62 |
+
| `prompt_deep` | `[L, P, 1, W]` | Optional; pooled into 5 prompt nodes for alignment |
|
| 63 |
+
|
| 64 |
+
| Output | Shape | Role |
|
| 65 |
+
|--------|-------|------|
|
| 66 |
+
| `root_feat` | `[B, D]` | Hierarchy-refined global representation |
|
| 67 |
+
| `mid_feats` | `[B, 2, D]` | Upper-half and lower-half mid nodes |
|
| 68 |
+
| `leaf_feats` | `[B, 4, D]` | Leaf node features (unchanged from input leaves) |
|
| 69 |
+
| `hier_loss` | scalar | Lorentz hierarchical contrastive loss |
|
| 70 |
+
| `prompt_loss` | scalar | Prompt tree + alignment loss (0 if disabled) |
|
| 71 |
+
|
| 72 |
+
Default feature dimension `D = 768` (CLIP ViT-L/14 projected dim).
|
| 73 |
+
|
| 74 |
+
---
|
| 75 |
+
|
| 76 |
+
## 2. Tree topology (depth 3)
|
| 77 |
+
|
| 78 |
+
Fixed binary tree — no learned tree search (anatomy is predefined):
|
| 79 |
+
|
| 80 |
+
```
|
| 81 |
+
root (whole entity)
|
| 82 |
+
/ \
|
| 83 |
+
mid_upper mid_lower
|
| 84 |
+
/ \ / \
|
| 85 |
+
leaf_0 leaf_1 leaf_2 leaf_3
|
| 86 |
+
(head) (upper) (lower) (feet)
|
| 87 |
+
|
| 88 |
+
region_tokens[:, 0] (global) ──► soft-anchors root via global_gate
|
| 89 |
+
global_feat ──► soft-anchors root via (1 - global_gate)
|
| 90 |
+
```
|
| 91 |
+
|
| 92 |
+
Parent table (`tree.py`):
|
| 93 |
+
|
| 94 |
+
- Leaves 0,1 → mid 0; leaves 2,3 → mid 1.
|
| 95 |
+
- Mids 0,1 → root.
|
| 96 |
+
|
| 97 |
+
---
|
| 98 |
+
|
| 99 |
+
## 3. Euclidean tree construction
|
| 100 |
+
|
| 101 |
+
### Leaves
|
| 102 |
+
|
| 103 |
+
```python
|
| 104 |
+
global_tok = region_tokens[:, 0] # [B, D]
|
| 105 |
+
leaves = region_tokens[:, 1:5] # [B, 4, D]
|
| 106 |
+
```
|
| 107 |
+
|
| 108 |
+
### Mid nodes — PairMerge
|
| 109 |
+
|
| 110 |
+
Two sibling leaves are merged with a small MLP:
|
| 111 |
+
|
| 112 |
+
```
|
| 113 |
+
PairMerge(a, b):
|
| 114 |
+
concat(a, b) → Linear(2D → D) → GELU → Linear(D → D)
|
| 115 |
+
```
|
| 116 |
+
|
| 117 |
+
```python
|
| 118 |
+
mid_upper = PairMerge(leaves[:, 0], leaves[:, 1])
|
| 119 |
+
mid_lower = PairMerge(leaves[:, 2], leaves[:, 3])
|
| 120 |
+
mid = stack([mid_upper, mid_lower], dim=1) # [B, 2, D]
|
| 121 |
+
```
|
| 122 |
+
|
| 123 |
+
### Root
|
| 124 |
+
|
| 125 |
+
```python
|
| 126 |
+
root_raw = PairMerge(mid[:, 0], mid[:, 1])
|
| 127 |
+
root = root_raw + gate * global_tok + (1 - gate) * global_feat
|
| 128 |
+
```
|
| 129 |
+
|
| 130 |
+
`gate` is a learnable scalar (`global_gate`, init 0.5).
|
| 131 |
+
|
| 132 |
+
---
|
| 133 |
+
|
| 134 |
+
## 4. CLIP Euclidean → Lorentz conversion block
|
| 135 |
+
|
| 136 |
+
**Code:** `lorentz.py` — `EuclideanToLorentz`, `expmap0_euclidean_space`
|
| 137 |
+
|
| 138 |
+
CLIP features (`all_class`, CLS, merged tree nodes) live in **Euclidean** space ℝ^D. Hierarchical losses need Lorentz points on the hyperboloid. Tree merges (`PairMerge`) stay Euclidean; conversion happens **after** the tree is built.
|
| 139 |
+
|
| 140 |
+
```
|
| 141 |
+
CLIP / tree nodes (Euclidean ℝ^D)
|
| 142 |
+
│
|
| 143 |
+
▼
|
| 144 |
+
┌─────────────────────────────────┐
|
| 145 |
+
│ EuclideanToLorentz │
|
| 146 |
+
│ Linear(D→D) adapter │
|
| 147 |
+
│ α · v (α learnable, init 1/��D) │
|
| 148 |
+
│ MERU expmap0_euclidean_space │
|
| 149 |
+
└─────────────────────────────────┘
|
| 150 |
+
│
|
| 151 |
+
▼
|
| 152 |
+
Lorentz points [time, space] ∈ ℝ^{D+1}
|
| 153 |
+
```
|
| 154 |
+
|
| 155 |
+
### MERU-style lift (not Minkowski-norm on CLIP vectors)
|
| 156 |
+
|
| 157 |
+
Treat Euclidean embedding `v ∈ ℝ^D` as **space components only** (tangent at the hyperboloid origin):
|
| 158 |
+
|
| 159 |
+
```
|
| 160 |
+
r = clamp(√c ‖v‖₂, max_norm)
|
| 161 |
+
x_space = unit(v) · sinh(r) / √c
|
| 162 |
+
x_time = sqrt(1/c + ‖x_space‖²)
|
| 163 |
+
h = [x_time, x_space] ∈ ℝ^{D+1}
|
| 164 |
+
```
|
| 165 |
+
|
| 166 |
+
Learnable scale `α` (init `1/√D`) is applied before the expmap to keep `sinh` numerically stable for CLIP-scale norms (MERU, ICML 2023).
|
| 167 |
+
|
| 168 |
+
### LorentzManifold
|
| 169 |
+
|
| 170 |
+
- Curvature `c` (default `1.0`, flag `--hi_mapper_curvature`).
|
| 171 |
+
- **`from_euclidean` / `EuclideanToLorentz`:** correct CLIP lift (above).
|
| 172 |
+
- **`geodesic_dist(x, y)`:** Lorentz distance via Minkowski inner product and `acosh` (on already-lifted points).
|
| 173 |
+
|
| 174 |
+
All tree nodes (leaves, mid, root) are mapped independently:
|
| 175 |
+
|
| 176 |
+
```
|
| 177 |
+
leaves_h [B, 4, D+1]
|
| 178 |
+
mid_h [B, 2, D+1]
|
| 179 |
+
root_h [B, 1, D+1]
|
| 180 |
+
```
|
| 181 |
+
|
| 182 |
+
Hyperbolic space is used because tree volume grows exponentially with depth; Euclidean embeddings distort those relations. Downstream classification still uses the **Euclidean** `root_feat`.
|
| 183 |
+
|
| 184 |
+
---
|
| 185 |
+
|
| 186 |
+
## 5. Hierarchical contrastive loss
|
| 187 |
+
|
| 188 |
+
**Code:** `tree.py` — `hierarchical_contrastive_loss`
|
| 189 |
+
|
| 190 |
+
Margin-based loss (`margin = 0.1`) on geodesic distances in Lorentz space:
|
| 191 |
+
|
| 192 |
+
| Constraint | Meaning |
|
| 193 |
+
|------------|---------|
|
| 194 |
+
| `d(leaf, parent_mid) < d(leaf, uncle_mid)` | Each leaf pulled toward its parent, not the sibling branch |
|
| 195 |
+
| `d(mid, root) < d(mid, sibling_mid)` | Mid nodes organized under root |
|
| 196 |
+
| `d(leaf, mid) < d(leaf, root)` | Leaves shallower than root (depth consistency) |
|
| 197 |
+
|
| 198 |
+
Formally, for each constraint:
|
| 199 |
+
|
| 200 |
+
```
|
| 201 |
+
loss += ReLU(d_parent - d_other + margin).mean()
|
| 202 |
+
```
|
| 203 |
+
|
| 204 |
+
This implements the Hi-Mapper idea: **child–parent similar, siblings/uncles dissimilar**, in a shared hyperbolic space.
|
| 205 |
+
|
| 206 |
+
---
|
| 207 |
+
|
| 208 |
+
## 6. Prompt-hierarchy alignment (optional)
|
| 209 |
+
|
| 210 |
+
**Flag:** `--optimize_prompts_with_hi_mapper`
|
| 211 |
+
**Weight:** `--prompt_hier_w` (default `0.05`)
|
| 212 |
+
|
| 213 |
+
When `prompt_deep` is provided:
|
| 214 |
+
|
| 215 |
+
1. **Pool** prompt tensor into 5 groups (same grouping as region tokens): mean over depth and token dims → `[5, D]`.
|
| 216 |
+
2. **Project** to feature dim if needed (`prompt_proj`: `W → D`).
|
| 217 |
+
3. **Build the same Lorentz tree** on prompt nodes.
|
| 218 |
+
4. **Losses:**
|
| 219 |
+
- `prompt_hier`: same hierarchical contrastive loss on the prompt tree.
|
| 220 |
+
- `align`: geodesic distance between prompt Lorentz nodes and **batch-mean** of visual Lorentz nodes (visual side detached).
|
| 221 |
+
|
| 222 |
+
```
|
| 223 |
+
prompt_loss = prompt_hier + align
|
| 224 |
+
```
|
| 225 |
+
|
| 226 |
+
This gives a direct training signal so auxiliary prompt parameters follow the same hierarchy as the visual region tokens.
|
| 227 |
+
|
| 228 |
+
---
|
| 229 |
+
|
| 230 |
+
## 7. HyperbolicDiffusionDecoder (generative aspect)
|
| 231 |
+
|
| 232 |
+
**Code:** `hyp_diffusion.py`
|
| 233 |
+
Inspired by **HypDAE** (ICCV 2025): the hierarchy acts as a semantic code; a small denoiser learns to reconstruct a global feature from that code.
|
| 234 |
+
|
| 235 |
+
### Hierarchy code
|
| 236 |
+
|
| 237 |
+
```python
|
| 238 |
+
z = concat(root_feat, mid_feats.flatten(1), leaf_feats.flatten(1))
|
| 239 |
+
# [B, D * (1 + 2 + 4)] = [B, 5376] when D=768
|
| 240 |
+
```
|
| 241 |
+
|
| 242 |
+
### DDPM step (feature space)
|
| 243 |
+
|
| 244 |
+
Training only — **not used at inference**.
|
| 245 |
+
|
| 246 |
+
1. Sample `t ~ Uniform{0, …, steps−1}`, `ε ~ N(0, I)`.
|
| 247 |
+
2. Noisy target: `x_t = √ᾱ_t · global_feat + √(1−ᾱ_t) · ε`.
|
| 248 |
+
3. Predict noise: `ε̂ = eps_θ(x_t, z, t_emb)`.
|
| 249 |
+
4. Loss: `MSE(ε̂, ε)`.
|
| 250 |
+
|
| 251 |
+
### Network
|
| 252 |
+
|
| 253 |
+
```
|
| 254 |
+
t_emb = MLP(t / steps) # [B, 128]
|
| 255 |
+
inp = [x_t ‖ z ‖ t_emb] # [B, D + 5376 + 128]
|
| 256 |
+
ε̂ = Linear → SiLU → Linear → SiLU → Linear
|
| 257 |
+
```
|
| 258 |
+
|
| 259 |
+
Default `steps = 6` (`--hyp_diffusion_steps`).
|
| 260 |
+
|
| 261 |
+
The decoder forces the hierarchy to be **generatively sufficient** for the global feature, without pixel-level image synthesis.
|
| 262 |
+
|
| 263 |
+
---
|
| 264 |
+
|
| 265 |
+
## 8. Module map
|
| 266 |
+
|
| 267 |
+
| File | Class / symbol | Role |
|
| 268 |
+
|------|----------------|------|
|
| 269 |
+
| `lorentz.py` | `EuclideanToLorentz`, `expmap0_euclidean_space`, `LorentzManifold` | CLIP Euclidean→Lorentz lift; geodesic distance; curvature |
|
| 270 |
+
| `tree.py` | `DIV_TREE`, `hierarchical_contrastive_loss`, `alignment_loss` | Fixed parent table and losses |
|
| 271 |
+
| `hi_mapper.py` | `DivHiMapper`, `PairMerge` | Euclidean tree build; Lorentz lift; prompt alignment |
|
| 272 |
+
| `hyp_diffusion.py` | `HyperbolicDiffusionDecoder` | CLS / global feature reconstruction DDPM |
|
| 273 |
+
|
| 274 |
+
---
|
| 275 |
+
|
| 276 |
+
## 9. Hyperparameters
|
| 277 |
+
|
| 278 |
+
| Parameter | Default | Description |
|
| 279 |
+
|-----------|---------|-------------|
|
| 280 |
+
| `feat_dim` (`D`) | 768 | Feature dimension |
|
| 281 |
+
| `div_num` | 4 | Number of region leaves (5 tokens total with global) |
|
| 282 |
+
| `hi_mapper_curvature` | 1.0 | Lorentz curvature `c` |
|
| 283 |
+
| `hi_mapper_w` | 0.1 | Weight on `hi_mapper_loss` in total training loss |
|
| 284 |
+
| `prompt_hier_w` | 0.05 | Weight on prompt alignment branch |
|
| 285 |
+
| `hyp_diffusion_w` | 0.1 | Weight on diffusion MSE |
|
| 286 |
+
| `hyp_diffusion_steps` | 6 | DDPM timesteps |
|
| 287 |
+
|
| 288 |
+
---
|
| 289 |
+
|
| 290 |
+
## 10. Training vs inference
|
| 291 |
+
|
| 292 |
+
| Component | Training | Inference |
|
| 293 |
+
|-----------|----------|-----------|
|
| 294 |
+
| DivHiMapper (tree + Lorentz) | ✓ | ✓ |
|
| 295 |
+
| `root_feat` output | ✓ | ✓ |
|
| 296 |
+
| Hierarchical contrastive loss | ✓ | — |
|
| 297 |
+
| Prompt alignment | optional | — |
|
| 298 |
+
| HyperbolicDiffusionDecoder | optional (aux loss) | **off** |
|
| 299 |
+
|
| 300 |
+
At inference, only the forward pass through DivHiMapper runs; no diffusion sampling, no extra loss terms.
|
| 301 |
+
|
| 302 |
+
---
|
| 303 |
+
|
| 304 |
+
## 11. Data flow (tensor shapes, D=768)
|
| 305 |
+
|
| 306 |
+
```
|
| 307 |
+
region_tokens [B,5,768]
|
| 308 |
+
global_feat [B,768]
|
| 309 |
+
│
|
| 310 |
+
├─ leaves [B,4,768]
|
| 311 |
+
├─ mid [B,2,768] ← PairMerge pairs
|
| 312 |
+
└─ root [B,768] ← PairMerge + gate blend
|
| 313 |
+
│
|
| 314 |
+
├─ EuclideanToLorentz → leaves_h [B,4,769], mid_h [B,2,769], root_h [B,1,769]
|
| 315 |
+
│
|
| 316 |
+
├─ hier_loss (scalar)
|
| 317 |
+
│
|
| 318 |
+
├─ prompt_loss (scalar, optional)
|
| 319 |
+
│
|
| 320 |
+
└─ z [B,5376] ──► diff_loss (scalar, optional, needs global_feat as target)
|
| 321 |
+
```
|
| 322 |
+
|
| 323 |
+
---
|
| 324 |
+
|
| 325 |
+
## 12. Design rationale
|
| 326 |
+
|
| 327 |
+
- **Fixed tree:** Region structure is known (4 body bands + global); no Mixture-of-Gaussians tree search as in the original Hi-Mapper paper.
|
| 328 |
+
- **Lorentz geometry:** Encodes exponential growth of hierarchy levels with lower distortion than Euclidean space.
|
| 329 |
+
- **Explicit Euclidean→Lorentz block:** CLIP features are Euclidean; MERU-style space-only `expmap0` (with learnable `α`) lifts them correctly onto the hyperboloid without treating the first CLIP coordinate as Minkowski time.
|
| 330 |
+
- **Generative decoder:** HypDAE-style auxiliary loss makes the hierarchy a predictive code for the global feature, improving root representation without generating images.
|
| 331 |
+
- **Prompt alignment:** Optional branch to synchronize learnable prompt parameters with the visual hierarchy in the same manifold.
|
| 332 |
+
|
| 333 |
+
---
|
| 334 |
+
|
| 335 |
+
## References
|
| 336 |
+
|
| 337 |
+
- Kwon et al., *Improving Visual Recognition with Hyperbolical Visual Hierarchy Mapping* (Hi-Mapper), CVPR 2024. [arXiv:2404.00974](https://arxiv.org/abs/2404.00974)
|
| 338 |
+
- Desai et al., *Hyperbolic Image-Text Representations* (MERU), ICML 2023. [arXiv:2304.09172](https://arxiv.org/abs/2304.09172)
|
| 339 |
+
- Li et al., *HypDAE: Hyperbolic Diffusion Autoencoders for Hierarchical Few-shot Image Generation*, ICCV 2025.
|
| 340 |
+
- Code: [kwonjunn01/Hi-Mapper](https://github.com/kwonjunn01/Hi-Mapper)
|
hi_mapper/__init__.py
ADDED
|
@@ -0,0 +1,26 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from .hi_mapper import DivHiMapper, build_attr_groups
|
| 2 |
+
from .hyp_diffusion import HyperbolicDiffusionDecoder
|
| 3 |
+
from .lorentz import (
|
| 4 |
+
EuclideanToLorentz,
|
| 5 |
+
LorentzManifold,
|
| 6 |
+
entailment_loss,
|
| 7 |
+
exterior_angle,
|
| 8 |
+
half_aperture,
|
| 9 |
+
)
|
| 10 |
+
from .tree import (
|
| 11 |
+
attribute_entailment_loss,
|
| 12 |
+
hierarchical_entailment_loss,
|
| 13 |
+
)
|
| 14 |
+
|
| 15 |
+
__all__ = [
|
| 16 |
+
"DivHiMapper",
|
| 17 |
+
"build_attr_groups",
|
| 18 |
+
"HyperbolicDiffusionDecoder",
|
| 19 |
+
"EuclideanToLorentz",
|
| 20 |
+
"LorentzManifold",
|
| 21 |
+
"entailment_loss",
|
| 22 |
+
"exterior_angle",
|
| 23 |
+
"half_aperture",
|
| 24 |
+
"attribute_entailment_loss",
|
| 25 |
+
"hierarchical_entailment_loss",
|
| 26 |
+
]
|
hi_mapper/hi_mapper.py
ADDED
|
@@ -0,0 +1,198 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Div-based Hi-Mapper with optional prompt-hierarchy alignment."""
|
| 2 |
+
from __future__ import annotations
|
| 3 |
+
|
| 4 |
+
import torch
|
| 5 |
+
import torch.nn as nn
|
| 6 |
+
|
| 7 |
+
from .lorentz import EuclideanToLorentz
|
| 8 |
+
from .tree import alignment_loss, attribute_entailment_loss, hierarchical_entailment_loss
|
| 9 |
+
|
| 10 |
+
# Node layout used when stacking the tree for attribute grounding:
|
| 11 |
+
# 0..3 leaves (head, upper body, lower body, feet), 4..5 mid, 6 root
|
| 12 |
+
LEAF_HEAD, LEAF_UPPER, LEAF_LOWER, LEAF_FEET = 0, 1, 2, 3
|
| 13 |
+
MID_UPPER, MID_LOWER, ROOT = 4, 5, 6
|
| 14 |
+
|
| 15 |
+
# PETA attribute-name prefixes mapped onto hierarchy nodes. Whole-person
|
| 16 |
+
# attributes (age, gender) sit at the root because they are the most general;
|
| 17 |
+
# carried objects span the torso, so they attach to the mid nodes; the rest are
|
| 18 |
+
# region-specific and attach to the matching row band.
|
| 19 |
+
_PREFIX_TO_NODE = (
|
| 20 |
+
("personal", ROOT),
|
| 21 |
+
("carrying", MID_UPPER),
|
| 22 |
+
("accessory", LEAF_HEAD),
|
| 23 |
+
("hair", LEAF_HEAD),
|
| 24 |
+
("head", LEAF_HEAD),
|
| 25 |
+
("upperbody", LEAF_UPPER),
|
| 26 |
+
("lowerbody", LEAF_LOWER),
|
| 27 |
+
("footwear", LEAF_FEET),
|
| 28 |
+
("shoes", LEAF_FEET),
|
| 29 |
+
)
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
def build_attr_groups(attr_names) -> dict[int, list[int]]:
|
| 33 |
+
"""
|
| 34 |
+
Map attribute indices onto hierarchy nodes by name prefix.
|
| 35 |
+
|
| 36 |
+
Returns ``{node_index: [attr_index, ...]}``. Attributes whose prefix is not
|
| 37 |
+
recognised are skipped rather than forced onto a node, so datasets without
|
| 38 |
+
PETA's naming convention simply produce fewer (or no) grounded constraints.
|
| 39 |
+
"""
|
| 40 |
+
groups: dict[int, list[int]] = {}
|
| 41 |
+
for idx, raw in enumerate(attr_names):
|
| 42 |
+
name = str(raw).strip().lower()
|
| 43 |
+
for prefix, node in _PREFIX_TO_NODE:
|
| 44 |
+
if name.startswith(prefix):
|
| 45 |
+
groups.setdefault(node, []).append(idx)
|
| 46 |
+
break
|
| 47 |
+
return groups
|
| 48 |
+
|
| 49 |
+
|
| 50 |
+
class PairMerge(nn.Module):
|
| 51 |
+
def __init__(self, dim: int):
|
| 52 |
+
super().__init__()
|
| 53 |
+
self.net = nn.Sequential(
|
| 54 |
+
nn.Linear(dim * 2, dim),
|
| 55 |
+
nn.GELU(),
|
| 56 |
+
nn.Linear(dim, dim),
|
| 57 |
+
)
|
| 58 |
+
|
| 59 |
+
def forward(self, a: torch.Tensor, b: torch.Tensor) -> torch.Tensor:
|
| 60 |
+
return self.net(torch.cat([a, b], dim=-1))
|
| 61 |
+
|
| 62 |
+
|
| 63 |
+
class DivHiMapper(nn.Module):
|
| 64 |
+
"""
|
| 65 |
+
Build a depth-3 tree from PromptPAR div tokens (all_class) and optional
|
| 66 |
+
prompt_deep slices, with a Lorentz entailment-cone hierarchy loss.
|
| 67 |
+
|
| 68 |
+
Tree merges stay in Euclidean space; CLIP Euclidean features are lifted
|
| 69 |
+
to the Lorentz hyperboloid via ``EuclideanToLorentz`` (MERU-style) before
|
| 70 |
+
the hierarchical losses.
|
| 71 |
+
"""
|
| 72 |
+
|
| 73 |
+
def __init__(
|
| 74 |
+
self,
|
| 75 |
+
feat_dim: int = 768,
|
| 76 |
+
prompt_width: int = 1024,
|
| 77 |
+
div_num: int = 4,
|
| 78 |
+
curvature: float = 1.0,
|
| 79 |
+
optimize_prompts: bool = False,
|
| 80 |
+
target_radius: float = 1.0,
|
| 81 |
+
entail_margin: float = 0.0,
|
| 82 |
+
attr_groups: dict[int, list[int]] | None = None,
|
| 83 |
+
):
|
| 84 |
+
super().__init__()
|
| 85 |
+
self.feat_dim = feat_dim
|
| 86 |
+
self.div_num = div_num
|
| 87 |
+
self.optimize_prompts = optimize_prompts
|
| 88 |
+
self.entail_margin = entail_margin
|
| 89 |
+
self.attr_groups = attr_groups or {}
|
| 90 |
+
|
| 91 |
+
self.euclid_to_lorentz = EuclideanToLorentz(
|
| 92 |
+
feat_dim, curvature=curvature, target_radius=target_radius
|
| 93 |
+
)
|
| 94 |
+
# Shared manifold reference for hierarchical / alignment losses
|
| 95 |
+
self.manifold = self.euclid_to_lorentz.manifold
|
| 96 |
+
|
| 97 |
+
# Text embeddings live on a different norm scale than visual tokens, so
|
| 98 |
+
# they get their own adapter and scale (MERU uses separate lambda_img /
|
| 99 |
+
# lambda_txt for exactly this reason) while sharing the manifold.
|
| 100 |
+
self.text_to_lorentz = None
|
| 101 |
+
if self.attr_groups:
|
| 102 |
+
self.text_to_lorentz = EuclideanToLorentz(
|
| 103 |
+
feat_dim, curvature=curvature, target_radius=target_radius
|
| 104 |
+
)
|
| 105 |
+
self.text_to_lorentz.manifold = self.manifold
|
| 106 |
+
self.prompt_proj = nn.Linear(prompt_width, feat_dim) if prompt_width != feat_dim else nn.Identity()
|
| 107 |
+
|
| 108 |
+
self.merge_pair = PairMerge(feat_dim)
|
| 109 |
+
self.merge_root = PairMerge(feat_dim)
|
| 110 |
+
self.global_gate = nn.Parameter(torch.tensor(0.5))
|
| 111 |
+
# Zero-init: at step 0 the root is exactly the CLIP CLS token, so
|
| 112 |
+
# enabling Hi-Mapper cannot damage the classifier before it has learned
|
| 113 |
+
# anything. Previously a randomly-initialised PairMerge output replaced
|
| 114 |
+
# CLS outright and cost ~10 mA points in epoch 1.
|
| 115 |
+
self.root_gate = nn.Parameter(torch.tensor(0.0))
|
| 116 |
+
|
| 117 |
+
def _to_lorentz(self, x: torch.Tensor) -> torch.Tensor:
|
| 118 |
+
"""Lift CLIP Euclidean features to Lorentz hyperboloid (MERU-style)."""
|
| 119 |
+
if x.dim() == 2:
|
| 120 |
+
x = x.unsqueeze(1)
|
| 121 |
+
return self.euclid_to_lorentz(x)
|
| 122 |
+
|
| 123 |
+
def _build_tree(self, tokens: torch.Tensor, cls_tok: torch.Tensor | None = None):
|
| 124 |
+
"""
|
| 125 |
+
tokens: [B, div_num+1, D] — index 0 global part token, 1..div_num row bands.
|
| 126 |
+
Returns euclidean root/mid/leaves and lorentz versions.
|
| 127 |
+
"""
|
| 128 |
+
assert tokens.shape[1] == self.div_num + 1, (
|
| 129 |
+
f"Expected {self.div_num + 1} div tokens, got {tokens.shape[1]}"
|
| 130 |
+
)
|
| 131 |
+
global_tok = tokens[:, 0]
|
| 132 |
+
leaves = tokens[:, 1:]
|
| 133 |
+
|
| 134 |
+
mid_upper = self.merge_pair(leaves[:, 0], leaves[:, 1])
|
| 135 |
+
mid_lower = self.merge_pair(leaves[:, 2], leaves[:, 3])
|
| 136 |
+
mid = torch.stack([mid_upper, mid_lower], dim=1)
|
| 137 |
+
|
| 138 |
+
delta = self.merge_root(mid[:, 0], mid[:, 1]) + self.global_gate * global_tok
|
| 139 |
+
if cls_tok is not None:
|
| 140 |
+
root = cls_tok + self.root_gate * delta
|
| 141 |
+
else:
|
| 142 |
+
root = delta
|
| 143 |
+
|
| 144 |
+
leaves_h = self._to_lorentz(leaves)
|
| 145 |
+
mid_h = self._to_lorentz(mid)
|
| 146 |
+
root_h = self._to_lorentz(root.unsqueeze(1))
|
| 147 |
+
|
| 148 |
+
hier_loss = hierarchical_entailment_loss(
|
| 149 |
+
leaves_h, mid_h, root_h, self.manifold, margin=self.entail_margin
|
| 150 |
+
)
|
| 151 |
+
return root, mid, leaves, leaves_h, mid_h, root_h, hier_loss
|
| 152 |
+
|
| 153 |
+
def _pool_prompt_groups(self, prompt_deep: torch.Tensor) -> torch.Tensor:
|
| 154 |
+
"""
|
| 155 |
+
Pool prompt_deep [vis_depth, prompt_num, 1, width] into [div_num+1, feat_dim].
|
| 156 |
+
Groups match div slices in clip/model.py (div_prompt_num = prompt_num // (div_num+1)).
|
| 157 |
+
"""
|
| 158 |
+
vis_depth, prompt_num, _, width = prompt_deep.shape
|
| 159 |
+
group_size = prompt_num // (self.div_num + 1)
|
| 160 |
+
nodes = []
|
| 161 |
+
for g in range(self.div_num + 1):
|
| 162 |
+
sl = prompt_deep[:, g * group_size : (g + 1) * group_size]
|
| 163 |
+
pooled = sl.mean(dim=(0, 1, 2)) # [width]
|
| 164 |
+
nodes.append(self.prompt_proj(pooled))
|
| 165 |
+
return torch.stack(nodes, dim=0)
|
| 166 |
+
|
| 167 |
+
def forward(
|
| 168 |
+
self,
|
| 169 |
+
all_class: torch.Tensor,
|
| 170 |
+
cls_tok: torch.Tensor,
|
| 171 |
+
prompt_deep: torch.Tensor | None = None,
|
| 172 |
+
text_features: torch.Tensor | None = None,
|
| 173 |
+
):
|
| 174 |
+
root, mid, leaves, leaves_h, mid_h, root_h, hier_loss = self._build_tree(all_class, cls_tok)
|
| 175 |
+
|
| 176 |
+
prompt_loss = torch.zeros((), device=all_class.device, dtype=hier_loss.dtype)
|
| 177 |
+
if self.optimize_prompts and prompt_deep is not None:
|
| 178 |
+
prompt_nodes = self._pool_prompt_groups(prompt_deep).unsqueeze(0) # [1, G, D]
|
| 179 |
+
prompt_nodes_b = prompt_nodes.expand(all_class.shape[0], -1, -1)
|
| 180 |
+
_, _, _, p_leaves_h, p_mid_h, p_root_h, prompt_hier = self._build_tree(
|
| 181 |
+
prompt_nodes_b, cls_tok=None
|
| 182 |
+
)
|
| 183 |
+
visual_stack = torch.cat([leaves_h, mid_h, root_h], dim=1)
|
| 184 |
+
prompt_stack = torch.cat([p_leaves_h, p_mid_h, p_root_h], dim=1)
|
| 185 |
+
align = alignment_loss(
|
| 186 |
+
prompt_stack, visual_stack.detach(), self.manifold, margin=self.entail_margin
|
| 187 |
+
)
|
| 188 |
+
prompt_loss = prompt_hier + align
|
| 189 |
+
|
| 190 |
+
attr_loss = torch.zeros((), device=all_class.device, dtype=hier_loss.dtype)
|
| 191 |
+
if self.attr_groups and text_features is not None and self.text_to_lorentz is not None:
|
| 192 |
+
node_stack = torch.cat([leaves_h, mid_h, root_h], dim=1) # [B, 7, D+1]
|
| 193 |
+
attr_h = self.text_to_lorentz(text_features.unsqueeze(1)).squeeze(1) # [A, D+1]
|
| 194 |
+
attr_loss = attribute_entailment_loss(
|
| 195 |
+
node_stack, attr_h, self.attr_groups, self.manifold, margin=self.entail_margin
|
| 196 |
+
)
|
| 197 |
+
|
| 198 |
+
return root, mid, leaves, hier_loss, prompt_loss, attr_loss
|
hi_mapper/hyp_diffusion.py
ADDED
|
@@ -0,0 +1,52 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""HypDAE-style feature-space diffusion decoder (auxiliary loss only)."""
|
| 2 |
+
from __future__ import annotations
|
| 3 |
+
|
| 4 |
+
import math
|
| 5 |
+
|
| 6 |
+
import torch
|
| 7 |
+
import torch.nn as nn
|
| 8 |
+
import torch.nn.functional as F
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
class HyperbolicDiffusionDecoder(nn.Module):
|
| 12 |
+
"""
|
| 13 |
+
Small DDPM-style denoiser in Euclidean feature space, conditioned on the
|
| 14 |
+
hierarchy code. Trained to reconstruct the CLS visual feature; not used
|
| 15 |
+
at inference time.
|
| 16 |
+
"""
|
| 17 |
+
|
| 18 |
+
def __init__(self, dim: int, code_dim: int, steps: int = 6, hidden: int = 512):
|
| 19 |
+
super().__init__()
|
| 20 |
+
self.dim = dim
|
| 21 |
+
self.steps = steps
|
| 22 |
+
betas = torch.linspace(1e-4, 2e-2, steps)
|
| 23 |
+
alphas = 1.0 - betas
|
| 24 |
+
alpha_bar = torch.cumprod(alphas, dim=0)
|
| 25 |
+
self.register_buffer("alpha_bar", alpha_bar)
|
| 26 |
+
|
| 27 |
+
t_dim = hidden // 4
|
| 28 |
+
self.time_mlp = nn.Sequential(
|
| 29 |
+
nn.Linear(1, t_dim),
|
| 30 |
+
nn.SiLU(),
|
| 31 |
+
nn.Linear(t_dim, t_dim),
|
| 32 |
+
)
|
| 33 |
+
self.eps_net = nn.Sequential(
|
| 34 |
+
nn.Linear(dim + code_dim + t_dim, hidden),
|
| 35 |
+
nn.SiLU(),
|
| 36 |
+
nn.Linear(hidden, hidden),
|
| 37 |
+
nn.SiLU(),
|
| 38 |
+
nn.Linear(hidden, dim),
|
| 39 |
+
)
|
| 40 |
+
|
| 41 |
+
def forward(self, hierarchy_code: torch.Tensor, cls_tok: torch.Tensor) -> torch.Tensor:
|
| 42 |
+
b = cls_tok.shape[0]
|
| 43 |
+
device = cls_tok.device
|
| 44 |
+
t = torch.randint(0, self.steps, (b,), device=device)
|
| 45 |
+
eps = torch.randn_like(cls_tok)
|
| 46 |
+
ab = self.alpha_bar[t].view(b, 1)
|
| 47 |
+
x_t = ab.sqrt() * cls_tok + (1.0 - ab).sqrt() * eps
|
| 48 |
+
|
| 49 |
+
t_emb = self.time_mlp((t.float() / max(self.steps - 1, 1)).unsqueeze(-1))
|
| 50 |
+
inp = torch.cat([x_t, hierarchy_code, t_emb], dim=-1)
|
| 51 |
+
eps_hat = self.eps_net(inp)
|
| 52 |
+
return F.mse_loss(eps_hat, eps)
|
hi_mapper/lorentz.py
ADDED
|
@@ -0,0 +1,320 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Self-contained Lorentz hyperboloid ops (adapted from Hi-Mapper / MERU, no geoopt)."""
|
| 2 |
+
from __future__ import annotations
|
| 3 |
+
|
| 4 |
+
import math
|
| 5 |
+
from typing import Union
|
| 6 |
+
|
| 7 |
+
import torch
|
| 8 |
+
import torch.nn as nn
|
| 9 |
+
|
| 10 |
+
# Tangent radius cap. Only guards against overflow: sinh(8) ~ 1.5e3 is safe in
|
| 11 |
+
# fp32. Anything near 1.0 silently discards feature magnitude and zeroes the
|
| 12 |
+
# gradient flowing through the radius.
|
| 13 |
+
EXP_MAX_NORM = 8.0
|
| 14 |
+
|
| 15 |
+
# Absolute bounds for the learnable curvature (log space).
|
| 16 |
+
MIN_CURVATURE = 0.01
|
| 17 |
+
MAX_CURVATURE = 10.0
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
def _eps(dtype: torch.dtype) -> float:
|
| 21 |
+
"""Precision floor appropriate to the dtype (fp64 tests need a tighter one)."""
|
| 22 |
+
return 1e-15 if dtype == torch.float64 else 1e-7
|
| 23 |
+
|
| 24 |
+
|
| 25 |
+
def _sqrt(x: torch.Tensor) -> torch.Tensor:
|
| 26 |
+
return torch.sqrt(torch.clamp_min(x, _eps(x.dtype) ** 2))
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
def acosh(x: torch.Tensor) -> torch.Tensor:
|
| 30 |
+
e = _eps(x.dtype)
|
| 31 |
+
x = torch.clamp_min(x, 1.0 + e)
|
| 32 |
+
z = torch.sqrt(torch.clamp_min(x.pow(2) - 1.0, e))
|
| 33 |
+
return torch.log(x + z)
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
def _inner(u: torch.Tensor, v: torch.Tensor, keepdim: bool = False, dim: int = -1) -> torch.Tensor:
|
| 37 |
+
"""
|
| 38 |
+
Minkowski inner product on ambient Lorentz coords [time, space].
|
| 39 |
+
|
| 40 |
+
Always evaluated in float64 and returned in float64. The expression
|
| 41 |
+
``-u_t·v_t + <u_s, v_s>`` cancels catastrophically once the hyperbolic
|
| 42 |
+
radius grows (the time component scales as cosh(r), so at r=7 the two terms
|
| 43 |
+
are ~1e6 apart from a result of order 1). In fp32 that destroys every digit
|
| 44 |
+
by r≈5; the surrounding ops are cheap at these tensor sizes, so the extra
|
| 45 |
+
precision costs effectively nothing.
|
| 46 |
+
"""
|
| 47 |
+
d = u.size(dim) - 1
|
| 48 |
+
uv = u.double() * v.double()
|
| 49 |
+
if not keepdim:
|
| 50 |
+
return -uv.narrow(dim, 0, 1).squeeze(dim) + uv.narrow(dim, 1, d).sum(dim=dim, keepdim=False)
|
| 51 |
+
return -uv.narrow(dim, 0, 1) + uv.narrow(dim, 1, d).sum(dim=dim, keepdim=True)
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
def _norm(u: torch.Tensor, keepdim: bool = False, dim: int = -1) -> torch.Tensor:
|
| 55 |
+
return _sqrt(_inner(u, u, keepdim=keepdim, dim=dim))
|
| 56 |
+
|
| 57 |
+
|
| 58 |
+
def project(x: torch.Tensor, k: Union[float, torch.Tensor], dim: int = -1, max_norm: float = EXP_MAX_NORM) -> torch.Tensor:
|
| 59 |
+
if not torch.is_tensor(k):
|
| 60 |
+
k = torch.tensor(k, device=x.device, dtype=x.dtype)
|
| 61 |
+
dn = x.size(dim) - 1
|
| 62 |
+
right = x.narrow(dim, 1, dn)
|
| 63 |
+
if max_norm:
|
| 64 |
+
right = torch.renorm(right, 2, dim if dim >= 0 else dim - x.dim(), max_norm)
|
| 65 |
+
left = _sqrt((1.0 / k) + (right * right).sum(dim=dim, keepdim=True))
|
| 66 |
+
return torch.cat((left, right), dim=dim)
|
| 67 |
+
|
| 68 |
+
|
| 69 |
+
def expmap0_euclidean_space(
|
| 70 |
+
v: torch.Tensor,
|
| 71 |
+
k: Union[float, torch.Tensor],
|
| 72 |
+
dim: int = -1,
|
| 73 |
+
max_norm: float = EXP_MAX_NORM,
|
| 74 |
+
) -> torch.Tensor:
|
| 75 |
+
"""
|
| 76 |
+
MERU-style expmap0 for Euclidean space vectors.
|
| 77 |
+
|
| 78 |
+
Treats ``v ∈ ℝ^D`` as space components only (tangent at the hyperboloid
|
| 79 |
+
origin). Does NOT apply Minkowski norm to CLIP Euclidean features.
|
| 80 |
+
|
| 81 |
+
x_space = sinh(√c ‖v‖) / (√c ‖v‖) · v
|
| 82 |
+
x_time = sqrt(1/c + ‖x_space‖²)
|
| 83 |
+
|
| 84 |
+
Returns Lorentz coords ``[..., D+1]`` with time first: ``[x_time, x_space]``.
|
| 85 |
+
"""
|
| 86 |
+
if not torch.is_tensor(k):
|
| 87 |
+
k = torch.tensor(k, device=v.device, dtype=v.dtype)
|
| 88 |
+
# Euclidean L2 norm of space components
|
| 89 |
+
v_norm = torch.linalg.vector_norm(v, dim=dim, keepdim=True).clamp_min(1e-8)
|
| 90 |
+
# Clamp tangent radius r = √c ‖v‖ for numerical stability (MERU / Hi-Mapper)
|
| 91 |
+
radius = (v_norm * _sqrt(k)).clamp_max(max_norm)
|
| 92 |
+
# x_space = sinh(√c ‖v‖) / (√c ‖v‖) · v ≡ unit(v) · sinh(r) / √c
|
| 93 |
+
v_unit = v / v_norm
|
| 94 |
+
x_space = v_unit * (torch.sinh(radius) / _sqrt(k))
|
| 95 |
+
x_time = _sqrt((1.0 / k) + (x_space * x_space).sum(dim=dim, keepdim=True))
|
| 96 |
+
return torch.cat((x_time, x_space), dim=dim)
|
| 97 |
+
|
| 98 |
+
|
| 99 |
+
def expmap0(u: torch.Tensor, k: Union[float, torch.Tensor], dim: int = -1) -> torch.Tensor:
|
| 100 |
+
"""
|
| 101 |
+
Exponential map at origin for ambient Lorentz tangent vectors [time, space].
|
| 102 |
+
|
| 103 |
+
Prefer ``expmap0_euclidean_space`` when lifting CLIP Euclidean features.
|
| 104 |
+
"""
|
| 105 |
+
if not torch.is_tensor(k):
|
| 106 |
+
k = torch.tensor(k, device=u.device, dtype=u.dtype)
|
| 107 |
+
nomin = _norm(u, keepdim=True, dim=dim).to(u.dtype)
|
| 108 |
+
safe = nomin.clamp_min(1e-8)
|
| 109 |
+
u_unit = u / safe
|
| 110 |
+
nomin = nomin.clamp_max(EXP_MAX_NORM)
|
| 111 |
+
l_v = torch.cosh(nomin)
|
| 112 |
+
r_v = torch.sinh(nomin) * u_unit
|
| 113 |
+
dn = r_v.size(dim) - 1
|
| 114 |
+
return torch.cat((l_v + r_v.narrow(dim, 0, 1), r_v.narrow(dim, 1, dn)), dim=dim)
|
| 115 |
+
|
| 116 |
+
|
| 117 |
+
def dist(x: torch.Tensor, y: torch.Tensor, k: Union[float, torch.Tensor], dim: int = -1) -> torch.Tensor:
|
| 118 |
+
"""
|
| 119 |
+
Lorentz geodesic distance, evaluated in the numerically stable form
|
| 120 |
+
|
| 121 |
+
d(x, y) = (2/√c) · asinh( √( c · ⟨x-y, x-y⟩_L / 4 ) )
|
| 122 |
+
|
| 123 |
+
which is algebraically identical to the textbook ``acosh(-c⟨x,y⟩_L)/√c``.
|
| 124 |
+
|
| 125 |
+
The textbook form is unusable for nearby points: ``-c⟨x,y⟩_L`` approaches 1
|
| 126 |
+
and ``acosh(1+ε) ≈ √(2ε)`` amplifies the square root of the coordinate
|
| 127 |
+
error, so fp32 inputs give d(x, x) ≈ 2e-3 instead of 0. Differencing first
|
| 128 |
+
avoids the cancellation entirely: when x == y the difference is exactly
|
| 129 |
+
zero, so the distance is exactly zero.
|
| 130 |
+
"""
|
| 131 |
+
if not torch.is_tensor(k):
|
| 132 |
+
k = torch.tensor(k, device=x.device, dtype=x.dtype)
|
| 133 |
+
diff = x.double() - y.double()
|
| 134 |
+
sq = _inner(diff, diff, dim=dim, keepdim=False).clamp_min(0.0)
|
| 135 |
+
kd = k.double()
|
| 136 |
+
return (2.0 * torch.asinh(_sqrt(kd * sq / 4.0)) / _sqrt(kd)).to(x.dtype)
|
| 137 |
+
|
| 138 |
+
|
| 139 |
+
def half_aperture(x: torch.Tensor, k: Union[float, torch.Tensor], min_radius: float = 0.1) -> torch.Tensor:
|
| 140 |
+
"""
|
| 141 |
+
Half-aperture of the entailment cone at ``x`` (MERU eq. 10, arbitrary curvature).
|
| 142 |
+
|
| 143 |
+
aper(x) = asin( 2K / (√c ‖x_space‖) ), K = min_radius = 0.1
|
| 144 |
+
|
| 145 |
+
The cone narrows as ``x`` moves away from the origin, so points near the
|
| 146 |
+
origin (general concepts) can entail many children.
|
| 147 |
+
"""
|
| 148 |
+
if not torch.is_tensor(k):
|
| 149 |
+
k = torch.tensor(k, device=x.device, dtype=x.dtype)
|
| 150 |
+
x_space_norm = torch.linalg.vector_norm(x[..., 1:], dim=-1).clamp_min(1e-8)
|
| 151 |
+
ratio = 2.0 * min_radius / (_sqrt(k) * x_space_norm)
|
| 152 |
+
return torch.asin(ratio.clamp(-1.0 + 1e-6, 1.0 - 1e-6))
|
| 153 |
+
|
| 154 |
+
|
| 155 |
+
def exterior_angle(x: torch.Tensor, y: torch.Tensor, k: Union[float, torch.Tensor]) -> torch.Tensor:
|
| 156 |
+
"""
|
| 157 |
+
Exterior angle ∠Oxy between the cone axis at ``x`` and the geodesic to ``y``
|
| 158 |
+
(MERU eq. 11, arbitrary curvature).
|
| 159 |
+
"""
|
| 160 |
+
if not torch.is_tensor(k):
|
| 161 |
+
k = torch.tensor(k, device=x.device, dtype=x.dtype)
|
| 162 |
+
cxy = k.double() * _inner(x, y, dim=-1, keepdim=False)
|
| 163 |
+
x_space_norm = torch.linalg.vector_norm(x[..., 1:].double(), dim=-1).clamp_min(1e-12)
|
| 164 |
+
num = y[..., 0].double() + x[..., 0].double() * cxy
|
| 165 |
+
den = x_space_norm * _sqrt(cxy.pow(2) - 1.0)
|
| 166 |
+
ratio = (num / den.clamp_min(1e-12)).clamp(-1.0 + 1e-9, 1.0 - 1e-9)
|
| 167 |
+
return torch.acos(ratio).to(x.dtype)
|
| 168 |
+
|
| 169 |
+
|
| 170 |
+
def entailment_loss(
|
| 171 |
+
parent: torch.Tensor,
|
| 172 |
+
child: torch.Tensor,
|
| 173 |
+
k: Union[float, torch.Tensor],
|
| 174 |
+
margin: float = 0.0,
|
| 175 |
+
) -> torch.Tensor:
|
| 176 |
+
"""
|
| 177 |
+
MERU entailment loss: penalize the child for lying outside the parent's cone.
|
| 178 |
+
|
| 179 |
+
L = max(0, ext(parent, child) - aper(parent) + margin)
|
| 180 |
+
|
| 181 |
+
``margin`` (HyCoCLIP, ICLR 2025) pushes children strictly inside the cone
|
| 182 |
+
rather than merely onto its boundary.
|
| 183 |
+
"""
|
| 184 |
+
ext = exterior_angle(parent, child, k)
|
| 185 |
+
aper = half_aperture(parent, k)
|
| 186 |
+
return torch.relu(ext - aper + margin)
|
| 187 |
+
|
| 188 |
+
|
| 189 |
+
class LorentzManifold(nn.Module):
|
| 190 |
+
"""Lorentz manifold with learnable curvature (log-space)."""
|
| 191 |
+
|
| 192 |
+
def __init__(self, curvature: float = 1.0, learnable: bool = True):
|
| 193 |
+
super().__init__()
|
| 194 |
+
self.log_k = nn.Parameter(torch.tensor(math.log(curvature)), requires_grad=learnable)
|
| 195 |
+
|
| 196 |
+
@property
|
| 197 |
+
def k(self) -> torch.Tensor:
|
| 198 |
+
return self.log_k.exp()
|
| 199 |
+
|
| 200 |
+
def clamp_k(self) -> None:
|
| 201 |
+
"""Clamp curvature to absolute bounds so it cannot drift to 0 or blow up."""
|
| 202 |
+
with torch.no_grad():
|
| 203 |
+
self.log_k.data.clamp_(math.log(MIN_CURVATURE), math.log(MAX_CURVATURE))
|
| 204 |
+
|
| 205 |
+
def to_hyperboloid(self, tangent: torch.Tensor) -> torch.Tensor:
|
| 206 |
+
"""Lift ambient Lorentz tangent [time, space] via expmap0."""
|
| 207 |
+
self.clamp_k()
|
| 208 |
+
return expmap0(tangent, self.k)
|
| 209 |
+
|
| 210 |
+
def from_euclidean(self, v: torch.Tensor) -> torch.Tensor:
|
| 211 |
+
"""Lift Euclidean space vectors (CLIP-style) via MERU expmap0."""
|
| 212 |
+
self.clamp_k()
|
| 213 |
+
return expmap0_euclidean_space(v, self.k)
|
| 214 |
+
|
| 215 |
+
def geodesic_dist(self, x: torch.Tensor, y: torch.Tensor) -> torch.Tensor:
|
| 216 |
+
self.clamp_k()
|
| 217 |
+
return dist(x, y, self.k)
|
| 218 |
+
|
| 219 |
+
def radius(self, x: torch.Tensor) -> torch.Tensor:
|
| 220 |
+
"""Geodesic distance from the origin, i.e. the point's 'depth'."""
|
| 221 |
+
self.clamp_k()
|
| 222 |
+
x_space_norm = torch.linalg.vector_norm(x[..., 1:], dim=-1)
|
| 223 |
+
return torch.asinh(_sqrt(self.k) * x_space_norm) / _sqrt(self.k)
|
| 224 |
+
|
| 225 |
+
def entailment(self, parent: torch.Tensor, child: torch.Tensor, margin: float = 0.0) -> torch.Tensor:
|
| 226 |
+
self.clamp_k()
|
| 227 |
+
return entailment_loss(parent, child, self.k, margin=margin)
|
| 228 |
+
|
| 229 |
+
|
| 230 |
+
class EuclideanToLorentz(nn.Module):
|
| 231 |
+
"""
|
| 232 |
+
Explicit CLIP Euclidean → Lorentz hyperboloid conversion (MERU-style).
|
| 233 |
+
|
| 234 |
+
Pipeline:
|
| 235 |
+
v_euclid → Linear(D→D) adapter → v / E[‖v‖] → α · v → expmap0_euclidean_space
|
| 236 |
+
|
| 237 |
+
The scale of the tangent vector is the whole ballgame. MERU's ``α = 1/√D``
|
| 238 |
+
assumes the encoder output has norm ≈ √D, which puts the tangent radius
|
| 239 |
+
``√c·α·‖v‖`` at ≈ 1. Raw CLIP ``all_class`` features have norm ≈ 8.3, not
|
| 240 |
+
√768 ≈ 27.7, and a default-init adapter shrank them a further ≈ 0.58x. The
|
| 241 |
+
result landed at radius ≈ 0.21, where ``sinh(r) ≈ r`` and the manifold is
|
| 242 |
+
numerically indistinguishable from Euclidean space.
|
| 243 |
+
|
| 244 |
+
The correction divides by a running estimate of the *dataset* mean norm
|
| 245 |
+
rather than normalising each sample. That distinction matters: per-sample
|
| 246 |
+
normalisation (LayerNorm, L2) forces every point onto the same radius, and
|
| 247 |
+
in hyperbolic space the radius *is* the hierarchy signal - general concepts
|
| 248 |
+
near the origin, specific ones near the boundary. Dividing by a shared
|
| 249 |
+
scalar fixes the global scale, which was the bug, while leaving the relative
|
| 250 |
+
magnitudes that encode depth intact.
|
| 251 |
+
|
| 252 |
+
``log_alpha`` is stored in log space (MERU; HyperVLM ICCVW 2025) so the scale
|
| 253 |
+
cannot collapse to zero or flip sign during training.
|
| 254 |
+
"""
|
| 255 |
+
|
| 256 |
+
def __init__(
|
| 257 |
+
self,
|
| 258 |
+
feat_dim: int,
|
| 259 |
+
curvature: float = 1.0,
|
| 260 |
+
learnable_curvature: bool = True,
|
| 261 |
+
use_adapter: bool = True,
|
| 262 |
+
max_norm: float = EXP_MAX_NORM,
|
| 263 |
+
target_radius: float = 1.0,
|
| 264 |
+
momentum: float = 0.05,
|
| 265 |
+
):
|
| 266 |
+
super().__init__()
|
| 267 |
+
self.feat_dim = feat_dim
|
| 268 |
+
self.max_norm = max_norm
|
| 269 |
+
self.target_radius = target_radius
|
| 270 |
+
self.momentum = momentum
|
| 271 |
+
self.manifold = LorentzManifold(curvature=curvature, learnable=learnable_curvature)
|
| 272 |
+
|
| 273 |
+
if use_adapter:
|
| 274 |
+
self.adapter = nn.Linear(feat_dim, feat_dim)
|
| 275 |
+
# Identity init: the adapter must not silently rescale the norm.
|
| 276 |
+
nn.init.eye_(self.adapter.weight)
|
| 277 |
+
nn.init.zeros_(self.adapter.bias)
|
| 278 |
+
else:
|
| 279 |
+
self.adapter = nn.Identity()
|
| 280 |
+
|
| 281 |
+
# EMA of the mean feature norm; -1 flags "not yet initialised" so the
|
| 282 |
+
# first batch seeds it exactly instead of dragging up from a guess.
|
| 283 |
+
self.register_buffer("running_norm", torch.tensor(-1.0))
|
| 284 |
+
|
| 285 |
+
# After the running-norm division the mean norm is 1, so
|
| 286 |
+
# α = target_radius / √c puts the *mean* radius at `target_radius`.
|
| 287 |
+
alpha0 = target_radius / math.sqrt(curvature)
|
| 288 |
+
self.log_alpha = nn.Parameter(torch.tensor(math.log(alpha0)))
|
| 289 |
+
|
| 290 |
+
@property
|
| 291 |
+
def alpha(self) -> torch.Tensor:
|
| 292 |
+
return self.log_alpha.exp()
|
| 293 |
+
|
| 294 |
+
def _scale(self, v: torch.Tensor) -> torch.Tensor:
|
| 295 |
+
norms = v.norm(dim=-1)
|
| 296 |
+
if self.training:
|
| 297 |
+
batch_mean = norms.mean().detach()
|
| 298 |
+
with torch.no_grad():
|
| 299 |
+
if self.running_norm.item() < 0:
|
| 300 |
+
self.running_norm.fill_(batch_mean.item())
|
| 301 |
+
else:
|
| 302 |
+
self.running_norm.mul_(1.0 - self.momentum).add_(self.momentum * batch_mean)
|
| 303 |
+
ref = self.running_norm if self.running_norm.item() > 0 else norms.mean().detach()
|
| 304 |
+
return v / ref.clamp_min(1e-6)
|
| 305 |
+
|
| 306 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 307 |
+
"""
|
| 308 |
+
Args:
|
| 309 |
+
x: Euclidean features ``[..., D]`` or ``[B, N, D]``.
|
| 310 |
+
Returns:
|
| 311 |
+
Lorentz points ``[..., D+1]`` with time first.
|
| 312 |
+
"""
|
| 313 |
+
*prefix, d = x.shape
|
| 314 |
+
assert d == self.feat_dim, f"Expected feat_dim={self.feat_dim}, got {d}"
|
| 315 |
+
# sinh/acosh are unstable in fp16; keep the whole lift in fp32.
|
| 316 |
+
flat = x.reshape(-1, d).float()
|
| 317 |
+
v = self.alpha * self._scale(self.adapter(flat))
|
| 318 |
+
self.manifold.clamp_k()
|
| 319 |
+
hyp = expmap0_euclidean_space(v, self.manifold.k, max_norm=self.max_norm)
|
| 320 |
+
return hyp.view(*prefix, d + 1)
|
hi_mapper/tree.py
ADDED
|
@@ -0,0 +1,165 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Fixed div hierarchy tree: 4 row-band leaves -> 2 mid -> 1 root."""
|
| 2 |
+
from __future__ import annotations
|
| 3 |
+
|
| 4 |
+
from dataclasses import dataclass
|
| 5 |
+
from typing import Tuple
|
| 6 |
+
|
| 7 |
+
import torch
|
| 8 |
+
import torch.nn.functional as F
|
| 9 |
+
|
| 10 |
+
from .lorentz import LorentzManifold
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
@dataclass(frozen=True)
|
| 14 |
+
class DivTree:
|
| 15 |
+
"""Parent indices for [leaf0..3, mid0..1, root]."""
|
| 16 |
+
|
| 17 |
+
depth: int = 3
|
| 18 |
+
num_leaves: int = 4
|
| 19 |
+
num_mid: int = 2
|
| 20 |
+
|
| 21 |
+
@property
|
| 22 |
+
def leaf_to_mid(self) -> Tuple[int, ...]:
|
| 23 |
+
return (0, 0, 1, 1)
|
| 24 |
+
|
| 25 |
+
@property
|
| 26 |
+
def mid_to_root(self) -> Tuple[int, ...]:
|
| 27 |
+
return (0, 0)
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
DIV_TREE = DivTree()
|
| 31 |
+
|
| 32 |
+
|
| 33 |
+
def hierarchical_entailment_loss(
|
| 34 |
+
leaves_h: torch.Tensor,
|
| 35 |
+
mid_h: torch.Tensor,
|
| 36 |
+
root_h: torch.Tensor,
|
| 37 |
+
manifold: LorentzManifold,
|
| 38 |
+
margin: float = 0.0,
|
| 39 |
+
sibling_margin: float = 0.2,
|
| 40 |
+
radius_margin: float = 0.1,
|
| 41 |
+
w_entail: float = 1.0,
|
| 42 |
+
w_sibling: float = 0.5,
|
| 43 |
+
w_radius: float = 0.5,
|
| 44 |
+
) -> torch.Tensor:
|
| 45 |
+
"""
|
| 46 |
+
Hyperbolic hierarchy loss on a depth-3 binary tree (MERU / HyCoCLIP style).
|
| 47 |
+
|
| 48 |
+
Three complementary terms:
|
| 49 |
+
|
| 50 |
+
1. **Entailment** - each child must lie inside its parent's entailment cone.
|
| 51 |
+
This is the partial-order constraint that actually encodes a hierarchy in
|
| 52 |
+
hyperbolic space.
|
| 53 |
+
2. **Sibling separation** - a leaf stays closer to its own parent than to its
|
| 54 |
+
uncle, keeping the two branches apart.
|
| 55 |
+
3. **Radius ordering** - root nearer the origin than mid, mid nearer than
|
| 56 |
+
leaves, so generality maps onto distance from the origin. Without this the
|
| 57 |
+
cones are free to sit at arbitrary depths and the hierarchy is unanchored.
|
| 58 |
+
|
| 59 |
+
Args:
|
| 60 |
+
leaves_h: [B, 4, D+1]
|
| 61 |
+
mid_h: [B, 2, D+1]
|
| 62 |
+
root_h: [B, 1, D+1]
|
| 63 |
+
"""
|
| 64 |
+
tree = DIV_TREE
|
| 65 |
+
root = root_h[:, 0]
|
| 66 |
+
|
| 67 |
+
# ---- 1. entailment: parent cone contains child -------------------------
|
| 68 |
+
entail = []
|
| 69 |
+
for i in range(tree.num_mid):
|
| 70 |
+
entail.append(manifold.entailment(root, mid_h[:, i], margin=margin))
|
| 71 |
+
for i in range(tree.num_leaves):
|
| 72 |
+
parent = mid_h[:, tree.leaf_to_mid[i]]
|
| 73 |
+
entail.append(manifold.entailment(parent, leaves_h[:, i], margin=margin))
|
| 74 |
+
entail_loss = torch.stack(entail, dim=0).mean()
|
| 75 |
+
|
| 76 |
+
# ---- 2. sibling separation --------------------------------------------
|
| 77 |
+
sibling = []
|
| 78 |
+
for i in range(tree.num_leaves):
|
| 79 |
+
mid_idx = tree.leaf_to_mid[i]
|
| 80 |
+
d_parent = manifold.geodesic_dist(leaves_h[:, i], mid_h[:, mid_idx])
|
| 81 |
+
d_uncle = manifold.geodesic_dist(leaves_h[:, i], mid_h[:, 1 - mid_idx])
|
| 82 |
+
sibling.append(F.relu(d_parent - d_uncle + sibling_margin))
|
| 83 |
+
sibling_loss = torch.stack(sibling, dim=0).mean()
|
| 84 |
+
|
| 85 |
+
# ---- 3. radius ordering: root < mid < leaf ------------------------------
|
| 86 |
+
r_root = manifold.radius(root)
|
| 87 |
+
r_mid = manifold.radius(mid_h)
|
| 88 |
+
r_leaf = manifold.radius(leaves_h)
|
| 89 |
+
|
| 90 |
+
radius = [F.relu(r_root.unsqueeze(-1) - r_mid + radius_margin).mean()]
|
| 91 |
+
for i in range(tree.num_leaves):
|
| 92 |
+
parent_r = r_mid[:, tree.leaf_to_mid[i]]
|
| 93 |
+
radius.append(F.relu(parent_r - r_leaf[:, i] + radius_margin).mean())
|
| 94 |
+
radius_loss = torch.stack(radius, dim=0).mean()
|
| 95 |
+
|
| 96 |
+
return w_entail * entail_loss + w_sibling * sibling_loss + w_radius * radius_loss
|
| 97 |
+
|
| 98 |
+
|
| 99 |
+
def alignment_loss(
|
| 100 |
+
prompt_h: torch.Tensor,
|
| 101 |
+
visual_h: torch.Tensor,
|
| 102 |
+
manifold: LorentzManifold,
|
| 103 |
+
margin: float = 0.0,
|
| 104 |
+
) -> torch.Tensor:
|
| 105 |
+
"""
|
| 106 |
+
Align prompt hierarchy nodes to the batch-mean visual hierarchy.
|
| 107 |
+
|
| 108 |
+
Uses entailment rather than raw geodesic distance: minimizing distance alone
|
| 109 |
+
drives every node onto the same point, which destroys the hierarchy it is
|
| 110 |
+
supposed to preserve. Here each visual node must entail its matching prompt
|
| 111 |
+
node, which pulls them together only along the cone axis.
|
| 112 |
+
"""
|
| 113 |
+
visual_mean = visual_h.mean(dim=0, keepdim=True)
|
| 114 |
+
if prompt_h.dim() == 2:
|
| 115 |
+
prompt_h = prompt_h.unsqueeze(0)
|
| 116 |
+
n = min(prompt_h.shape[1], visual_mean.shape[1])
|
| 117 |
+
prompt_nodes = prompt_h[:, :n]
|
| 118 |
+
visual_nodes = visual_mean[:, :n]
|
| 119 |
+
if prompt_nodes.shape[0] != visual_nodes.shape[0]:
|
| 120 |
+
visual_nodes = visual_nodes.expand(prompt_nodes.shape[0], -1, -1)
|
| 121 |
+
return manifold.entailment(visual_nodes, prompt_nodes, margin=margin).mean()
|
| 122 |
+
|
| 123 |
+
|
| 124 |
+
def attribute_entailment_loss(
|
| 125 |
+
node_h: torch.Tensor,
|
| 126 |
+
attr_h: torch.Tensor,
|
| 127 |
+
groups: dict[int, list[int]],
|
| 128 |
+
manifold: LorentzManifold,
|
| 129 |
+
margin: float = 0.0,
|
| 130 |
+
) -> torch.Tensor:
|
| 131 |
+
"""
|
| 132 |
+
Ground the visual hierarchy in the attribute label semantics (HyCoCLIP-style
|
| 133 |
+
compositional entailment).
|
| 134 |
+
|
| 135 |
+
Each hierarchy node must entail the text embeddings of the attributes that
|
| 136 |
+
belong to it, e.g. the root entails the whole-person attributes (age,
|
| 137 |
+
gender) while the lower-body leaf entails ``lowerBody*``. This gives the
|
| 138 |
+
hyperbolic geometry a label-grounded objective instead of a purely
|
| 139 |
+
self-supervised one.
|
| 140 |
+
|
| 141 |
+
Args:
|
| 142 |
+
node_h: [B, N, D+1] hierarchy nodes in Lorentz coords.
|
| 143 |
+
attr_h: [A, D+1] attribute text embeddings in Lorentz coords.
|
| 144 |
+
groups: node index -> list of attribute indices.
|
| 145 |
+
"""
|
| 146 |
+
if not groups:
|
| 147 |
+
return node_h.new_zeros(())
|
| 148 |
+
|
| 149 |
+
losses = []
|
| 150 |
+
b = node_h.shape[0]
|
| 151 |
+
for node_idx, attr_ids in groups.items():
|
| 152 |
+
if not attr_ids or node_idx >= node_h.shape[1]:
|
| 153 |
+
continue
|
| 154 |
+
parent = node_h[:, node_idx].unsqueeze(1).expand(-1, len(attr_ids), -1)
|
| 155 |
+
child = attr_h[attr_ids].unsqueeze(0).expand(b, -1, -1)
|
| 156 |
+
losses.append(manifold.entailment(parent, child, margin=margin).mean())
|
| 157 |
+
|
| 158 |
+
if not losses:
|
| 159 |
+
return node_h.new_zeros(())
|
| 160 |
+
return torch.stack(losses, dim=0).mean()
|
| 161 |
+
|
| 162 |
+
|
| 163 |
+
# Backwards-compatible alias; the margin-only formulation it used to implement
|
| 164 |
+
# could not express a hierarchy (no radius ordering, no cones).
|
| 165 |
+
hierarchical_contrastive_loss = hierarchical_entailment_loss
|
requirements.txt
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
torch>=2.0.0
|