ZACK777 commited on
Commit
617d016
·
verified ·
1 Parent(s): 51b22dd

Add HI-Mapper source (code only, no weights)

Browse files
.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
- *.rar filter=lfs diff=lfs merge=lfs -text
25
  *.safetensors filter=lfs diff=lfs merge=lfs -text
26
- saved_model/**/* filter=lfs diff=lfs merge=lfs -text
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