summaryrefslogtreecommitdiff
path: root/worldalign/io.py
diff options
context:
space:
mode:
Diffstat (limited to 'worldalign/io.py')
-rw-r--r--worldalign/io.py30
1 files changed, 30 insertions, 0 deletions
diff --git a/worldalign/io.py b/worldalign/io.py
new file mode 100644
index 0000000..6cac007
--- /dev/null
+++ b/worldalign/io.py
@@ -0,0 +1,30 @@
+from __future__ import annotations
+
+import torch
+
+from .common import normalized
+
+
+def load_feature_pair(
+ vision_path: str, text_path: str
+) -> tuple[dict, dict, dict[int, int], dict[int, int]]:
+ vision = torch.load(vision_path, map_location="cpu", weights_only=False)
+ text = torch.load(text_path, map_location="cpu", weights_only=False)
+ vision["features"] = normalized(vision["features"])
+ text["features"] = normalized(text["features"])
+ vision_lookup = {int(row): i for i, row in enumerate(vision["rows"])}
+ text_lookup = {int(row): i for i, row in enumerate(text["rows"])}
+ return vision, text, vision_lookup, text_lookup
+
+
+def select_rows(
+ feature: torch.Tensor, lookup: dict[int, int], rows: list[int]
+) -> torch.Tensor:
+ missing = [int(row) for row in rows if int(row) not in lookup]
+ if missing:
+ raise KeyError(
+ f"{len(missing)} requested rows are absent from feature cache; "
+ f"first missing rows: {missing[:5]}"
+ )
+ return feature[[lookup[int(row)] for row in rows]]
+