cozyclay 1.9.0 → 2.0.0
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- package/CHANGELOG.md +171 -0
- package/bin/agent/agent-routes.mjs +74 -70
- package/bin/agent/agent-runner.mjs +10 -3
- package/bin/agent/motion-runtime.mjs +32 -260
- package/bin/agent/providers.mjs +46 -3
- package/bin/agent/session-store.mjs +41 -5
- package/bin/agent/studio-prompt.mjs +1 -1
- package/bin/agent/studio-tools.mjs +48 -11
- package/bin/cozyclay.mjs +26 -8
- package/bin/live/cli.mjs +22 -3
- package/bin/mcp-runtime.mjs +35 -20
- package/bin/telemetry-state.mjs +5 -2
- package/bin/update-check.mjs +15 -0
- package/dist/app/index.html +6 -6
- package/dist/assets/analytics-MVZOEW8O.js +1 -0
- package/dist/assets/app-C1sdVbg4.css +1 -0
- package/dist/assets/app-CA55tcI3.js +4876 -0
- package/dist/assets/{demo-B-eg1CT8.js → demo-ByeYp-dY.js} +1 -1
- package/dist/assets/{first-shot-handoff-CGgnvNu8.js → first-shot-handoff-DgT28ZiV.js} +1 -1
- package/dist/assets/{landing-IbRp3FnQ.js → landing-BaXkpFip.js} +1 -1
- package/dist/assets/shot-prompt-CbqWcWYi.css +1 -0
- package/dist/assets/shot-prompt-Yrt99wZx.js +17 -0
- package/dist/assets/{ticket-BgIzOBOf.js → ticket-BcmsnHn5.js} +1 -1
- package/dist/assets/workflow-DRdwoDBT.css +1 -0
- package/dist/assets/workflow-dNPGY9xz.js +179 -0
- package/dist/cozyclay-package.json +1 -1
- package/dist/fonts/IBM-Plex-Sans-OFL.txt +92 -0
- package/dist/fonts/JetBrains-Mono-OFL.txt +93 -0
- package/dist/fonts/README.md +13 -9
- package/dist/fonts/ibm-plex-sans-400-latin.woff2 +0 -0
- package/dist/fonts/ibm-plex-sans-500-latin.woff2 +0 -0
- package/dist/fonts/ibm-plex-sans-600-latin.woff2 +0 -0
- package/dist/fonts/jetbrains-mono-400-latin.woff2 +0 -0
- package/dist/fonts/jetbrains-mono-500-latin.woff2 +0 -0
- package/dist/index.html +7 -7
- package/dist/privacy/index.html +13 -8
- package/dist/sitemap.xml +6 -6
- package/dist/workflow/index.html +5 -5
- package/mcp/LIVE-PROTOCOL.md +1 -1
- package/mcp/live-hub.mjs +35 -93
- package/mcp/mesh-file.mjs +48 -0
- package/mcp/server.mjs +6 -77
- package/mcp/tool-handlers.mjs +224 -377
- package/package.json +2 -1
- package/src/App.jsx +1273 -9075
- package/src/analytics.js +393 -9
- package/src/app-context.js +211 -0
- package/src/app-stage.jsx +25 -24
- package/src/ardy/auto-fix-panel.css +108 -0
- package/src/ardy/collision-blockers.js +9 -3
- package/src/ardy/fix-collisions.js +10 -66
- package/src/ardy/ground.js +4 -4
- package/src/ardy/ik-drag.js +101 -0
- package/src/ardy/ik-key-json.js +43 -0
- package/src/ardy/ik.js +122 -76
- package/src/ardy/physics-panel.css +2 -0
- package/src/ardy/physics-panel.jsx +17 -5
- package/src/ardy/platform-fit-panel.css +64 -0
- package/src/ardy/platform-fit-panel.jsx +40 -0
- package/src/ardy/platform-fit.js +326 -0
- package/src/ardy/playback.js +166 -20
- package/src/ardy/range-pin.js +299 -0
- package/src/ardy/timeline-coordinates.js +26 -0
- package/src/ardy/timeline.css +1199 -0
- package/src/ardy/timeline.jsx +149 -41
- package/src/ardy/waypoints.js +2 -2
- package/src/asset-pane.css +462 -0
- package/src/asset-pane.jsx +461 -232
- package/src/command-bus.js +349 -0
- package/src/commands/ai.js +31 -0
- package/src/commands/cast.js +140 -0
- package/src/commands/elements/character.js +13 -0
- package/src/commands/elements/motion.js +14 -0
- package/src/commands/elements/object.js +22 -0
- package/src/commands/elements/scene.js +16 -0
- package/src/commands/elements/shot.js +9 -0
- package/src/commands/elements/stage.js +11 -0
- package/src/commands/elements.js +155 -0
- package/src/commands/export.js +20 -0
- package/src/commands/index.js +40 -0
- package/src/commands/motion.js +231 -0
- package/src/commands/objects.js +174 -0
- package/src/commands/project.js +53 -0
- package/src/commands/scene.js +81 -0
- package/src/commands/shared.js +19 -0
- package/src/commands/shot.js +103 -0
- package/src/commands/stage.js +17 -0
- package/src/commands/view.js +42 -0
- package/src/document-store.js +144 -0
- package/src/domains/cast.js +1007 -0
- package/src/domains/motion.js +3219 -0
- package/src/domains/objects.js +932 -0
- package/src/domains/scenes.js +817 -0
- package/src/domains/shots.js +404 -0
- package/src/domains/stage.js +93 -0
- package/src/dualview.jsx +113 -57
- package/src/facing-marks.js +3 -0
- package/src/first-success-guide.jsx +14 -14
- package/src/grid-view.js +18 -5
- package/src/hierarchy-model.js +24 -16
- package/src/hierarchy-panel.css +386 -0
- package/src/hierarchy-panel.jsx +150 -21
- package/src/{fal-motion-client.js → i2v-motion-client.js} +13 -13
- package/src/i2v-motion-studio.jsx +180 -0
- package/src/ik-camera.js +3 -0
- package/src/main.jsx +6 -0
- package/src/motion/generation.js +105 -0
- package/src/motion-readiness-ui.jsx +4 -4
- package/src/motion-readiness.js +11 -0
- package/src/motion-trail.js +366 -9
- package/src/object-gizmo.jsx +7 -1
- package/src/otio.js +2 -1
- package/src/panels/CameraPanel.jsx +60 -0
- package/src/panels/CharacterTransformPanel.jsx +48 -0
- package/src/panels/EnvironmentPanel.jsx +40 -0
- package/src/panels/Foldout.jsx +32 -0
- package/src/panels/LightPanel.jsx +24 -0
- package/src/panels/ObjectTransformPanel.jsx +554 -0
- package/src/panels/PosePanel.jsx +97 -0
- package/src/panels/ProjectPanel.jsx +33 -0
- package/src/panels/PromptBlocksPanel.jsx +442 -0
- package/src/panels/PropsPanel.jsx +70 -0
- package/src/panels/ReferenceImageField.jsx +91 -0
- package/src/panels/RigControlPanel.jsx +78 -0
- package/src/panels/RigPanel.jsx +37 -0
- package/src/panels/SubjectBox.jsx +47 -0
- package/src/panels/SubjectsPanel.jsx +39 -0
- package/src/panels/VideoCapturePanel.jsx +180 -0
- package/src/panels/details.css +1171 -0
- package/src/panels/motion.css +197 -0
- package/src/panels/pose.css +39 -0
- package/src/planview.jsx +60 -31
- package/src/posestudio.jsx +59 -4
- package/src/project-browser.css +1073 -0
- package/src/project-browser.jsx +284 -123
- package/src/range-pin-object-transform.js +19 -0
- package/src/range-pin-panel.css +518 -0
- package/src/range-pin-panel.jsx +344 -0
- package/src/result-modal.jsx +5 -5
- package/src/room.jsx +79 -51
- package/src/scene-objects.js +37 -1
- package/src/scenes.js +5 -0
- package/src/semantic-edit.js +2 -0
- package/src/settings-menu.jsx +39 -105
- package/src/shell/BottomDock.jsx +450 -0
- package/src/shell/DetailsSlot.jsx +513 -0
- package/src/shell/LibrarySlot.jsx +86 -0
- package/src/shell/MenuBar.jsx +601 -0
- package/src/shell/OutlinerSlot.jsx +50 -0
- package/src/shell/PreferencesDialog.jsx +538 -0
- package/src/shell/PreferencesSlot.jsx +40 -0
- package/src/shell/StatusBar.jsx +125 -0
- package/src/shell/StudioShell.jsx +92 -0
- package/src/shell/TopBar.jsx +129 -0
- package/src/shell/ViewportToolbar.jsx +656 -0
- package/src/shell/agent-glass.css +296 -0
- package/src/shell/dock.css +93 -0
- package/src/shell/glass-regions.css +985 -0
- package/src/shell/glass.css +687 -0
- package/src/shell/log-store.js +52 -0
- package/src/shell/mode.css +9 -0
- package/src/shell/preferences.css +651 -0
- package/src/shell/shell.css +282 -0
- package/src/shell/studio-shell-context.js +19 -0
- package/src/shell/topbar.css +435 -0
- package/src/shell/viewport.css +451 -0
- package/src/store/authored-intent.js +22 -0
- package/src/store/runtime-adapters.js +92 -0
- package/src/store/scene-stage.js +21 -0
- package/src/store/use-document-store.js +14 -0
- package/src/studio-actions.js +210 -0
- package/src/studio-agent-commands.js +41 -271
- package/src/studio-agent-context.js +39 -10
- package/src/studio-agent-motion.js +122 -373
- package/src/studio-agent-protocol.js +67 -31
- package/src/studio-app-binding.js +379 -0
- package/src/studio-contact-sheet.js +73 -0
- package/src/studio-elements.js +58 -51
- package/src/styles/themes.css +200 -0
- package/src/styles/tokens.css +69 -0
- package/src/styles.css +460 -1742
- package/src/theme.js +44 -0
- package/src/timeline-extent.js +16 -0
- package/src/trail-key-conflicts.js +39 -0
- package/src/trail-pick.js +59 -0
- package/src/ui.jsx +8 -8
- package/src/use-case-question.jsx +66 -0
- package/src/workflow/AgentPanel.jsx +82 -69
- package/src/workflow/agent-client.js +12 -2
- package/src/workflow/agent-panel.css +34 -32
- package/src/workflow/cozy-scene-node.css +2 -0
- package/src/workflow/workflow.css +2 -0
- package/tools/ardy/__pycache__/cclay_gvhmr_worker.cpython-313.pyc +0 -0
- package/tools/ardy/bridge.mjs +164 -146
- package/tools/ardy/visual-qa.mjs +5 -5
- package/tools/bench/EXP3.md +61 -0
- package/tools/bench/cclay_bench_extract_incam.py +125 -0
- package/tools/bench/cclay_bench_extract_obs.py +204 -0
- package/tools/bench/cclay_bench_runner.py +42 -0
- package/tools/bench/cube-contact.mjs +162 -0
- package/tools/bench/exp3.mjs +152 -0
- package/tools/bench/extract-bench-lib.mjs +216 -0
- package/tools/bench/extract-bench.mjs +302 -0
- package/tools/bench/fal-generate.mjs +88 -0
- package/tools/bench/fit/README.md +189 -0
- package/tools/bench/fit/camera.mjs +59 -0
- package/tools/bench/fit/contact.mjs +419 -0
- package/tools/bench/fit/footlock.mjs +234 -0
- package/tools/bench/fit/motion.mjs +85 -0
- package/tools/bench/fit/pin.mjs +27 -0
- package/tools/bench/fit/remote.mjs +105 -0
- package/tools/bench/fit-bench.mjs +126 -0
- package/tools/bench/fit-sanity.mjs +45 -0
- package/tools/bench/metrics.mjs +313 -0
- package/tools/bench/obs/depth.mjs +258 -0
- package/tools/bench/obs/extrinsics.mjs +175 -0
- package/tools/bench/obs/fit_mannequin_betas.py +340 -0
- package/tools/bench/obs/ground.mjs +294 -0
- package/tools/bench/obs/heading.mjs +151 -0
- package/tools/bench/obs/ladder.mjs +555 -0
- package/tools/bench/obs/mannequin-betas.json +124 -0
- package/tools/bench/obs/remote.mjs +83 -0
- package/tools/bench/obs/rest_joints.py +52 -0
- package/tools/bench/obs/ybot-targets.mjs +49 -0
- package/tools/bench/obs-bench.mjs +520 -0
- package/tools/bench/score.mjs +343 -0
- package/tools/bench/summarize.mjs +86 -0
- package/tools/dev/pages/privacy.html +13 -8
- package/tools/gt-render/browser.mjs +159 -0
- package/tools/gt-render/camera-math.mjs +276 -0
- package/tools/gt-render/page.mjs +267 -0
- package/tools/gt-render/render.mjs +514 -0
- package/tools/gt-render/scene-box.mjs +31 -0
- package/tools/gt-render/take-transform.mjs +87 -0
- package/tools/morphgs/assets/gen_truth.py +39 -0
- package/tools/morphgs/assets/mesh_ori_rig.txt +27 -0
- package/tools/morphgs/assets/playback-check.mjs +30 -0
- package/tools/morphgs/demo-gate.sh +26 -0
- package/tools/morphgs/fbx2morphgs.mjs +68 -0
- package/tools/morphgs/morphgs-to-cskel27.mjs +105 -0
- package/tools/morphgs/patches/preprocess_src-none-mode.patch +76 -0
- package/tools/morphgs/setup-on-cluster.sh +126 -0
- package/tools/qa/css-rule-usage.mjs +283 -0
- package/tools/qa/studio-control-count.mjs +70 -32
- package/tools/run-tests.mjs +134 -12
- package/tools/track/DESIGN.md +513 -0
- package/tools/track/backfill-provenance.mjs +116 -0
- package/tools/track/budget.mjs +8 -0
- package/tools/track/check-rig.mjs +209 -0
- package/tools/track/diagnostics.schema.json +60 -0
- package/tools/track/export-rig.mjs +256 -0
- package/tools/track/fallback.mjs +7 -0
- package/tools/track/fk-parity-fixture.mjs +161 -0
- package/tools/track/gate.mjs +349 -0
- package/tools/track/masks.mjs +318 -0
- package/tools/track/metrics.mjs +201 -0
- package/tools/track/publish-obs.mjs +108 -0
- package/tools/track/py/check_env.py +61 -0
- package/tools/track/py/eval_lr.py +229 -0
- package/tools/track/py/lr_viterbi.py +294 -0
- package/tools/track/py/masks.py +405 -0
- package/tools/track/py/objective.py +382 -0
- package/tools/track/py/rig.py +508 -0
- package/tools/track/py/scene.py +326 -0
- package/tools/track/py/test_joint_indices.py +123 -0
- package/tools/track/py/test_lr_viterbi.py +195 -0
- package/tools/track/py/test_masks.py +229 -0
- package/tools/track/py/test_rig.py +288 -0
- package/tools/track/py/test_scene.py +354 -0
- package/tools/track/py/test_track.py +328 -0
- package/tools/track/py/track.py +411 -0
- package/tools/track/remote.mjs +140 -0
- package/tools/track/rig-dump.mjs +227 -0
- package/tools/track/run-box-tests.mjs +14 -0
- package/tools/track/run-box.mjs +168 -0
- package/tools/track/setup-box.sh +59 -0
- package/tools/track/study-2d.mjs +728 -0
- package/dist/assets/analytics-B1hnH66c.js +0 -1
- package/dist/assets/app-BWusbqgO.js +0 -4861
- package/dist/assets/app-qKDo4PBX.css +0 -1
- package/dist/assets/shot-prompt-CHUtw6af.js +0 -17
- package/dist/assets/shot-prompt-C_g2BHVc.css +0 -1
- package/dist/assets/workflow-CVwzMhz2.js +0 -179
- package/dist/assets/workflow-DsudxKHF.css +0 -1
- package/src/fal-motion-studio.jsx +0 -180
- package/src/scene-history.js +0 -129
|
@@ -0,0 +1,61 @@
|
|
|
1
|
+
#!/usr/bin/env python3
|
|
2
|
+
"""Small, real CUDA/SAM2 smoke test for the CozyFit box environment."""
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import importlib
|
|
6
|
+
import importlib.metadata
|
|
7
|
+
import os
|
|
8
|
+
import sys
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
def version(module_name: str, distribution: str) -> str:
|
|
12
|
+
importlib.import_module(module_name)
|
|
13
|
+
return importlib.metadata.version(distribution)
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
def main() -> int:
|
|
17
|
+
import numpy as np
|
|
18
|
+
import torch
|
|
19
|
+
|
|
20
|
+
print(f"python {sys.version.split()[0]}")
|
|
21
|
+
print(f"torch {version('torch', 'torch')}")
|
|
22
|
+
print(f"torchvision {version('torchvision', 'torchvision')}")
|
|
23
|
+
print(f"numpy {version('numpy', 'numpy')}")
|
|
24
|
+
print(f"scipy {version('scipy', 'scipy')}")
|
|
25
|
+
print(f"opencv {version('cv2', 'opencv-python-headless')}")
|
|
26
|
+
print(f"sam2 {version('sam2', 'SAM-2')}")
|
|
27
|
+
print(f"fast-simplification {version('fast_simplification', 'fast-simplification')}")
|
|
28
|
+
print(f"pytest {version('pytest', 'pytest')}")
|
|
29
|
+
cuda = bool(torch.cuda.is_available())
|
|
30
|
+
print(f"torch.cuda.is_available() {cuda}")
|
|
31
|
+
if not cuda:
|
|
32
|
+
raise RuntimeError("CUDA is unavailable")
|
|
33
|
+
|
|
34
|
+
from sam2.build_sam import build_sam2
|
|
35
|
+
from sam2.sam2_image_predictor import SAM2ImagePredictor
|
|
36
|
+
|
|
37
|
+
root = os.path.expanduser("~/cclay-ingest/cozyfit")
|
|
38
|
+
checkpoint = os.path.join(root, "checkpoints", "sam2.1_hiera_small.pt")
|
|
39
|
+
config = "configs/sam2.1/sam2.1_hiera_s.yaml"
|
|
40
|
+
torch.cuda.reset_peak_memory_stats()
|
|
41
|
+
model = build_sam2(config, checkpoint, device="cuda")
|
|
42
|
+
predictor = SAM2ImagePredictor(model)
|
|
43
|
+
frame = np.zeros((480, 832, 3), dtype=np.uint8)
|
|
44
|
+
predictor.set_image(frame)
|
|
45
|
+
predictor.predict(
|
|
46
|
+
point_coords=np.array([[416.0, 240.0]], dtype=np.float32),
|
|
47
|
+
point_labels=np.array([1], dtype=np.int32),
|
|
48
|
+
box=np.array([208.0, 120.0, 624.0, 360.0], dtype=np.float32),
|
|
49
|
+
multimask_output=False,
|
|
50
|
+
)
|
|
51
|
+
torch.cuda.synchronize()
|
|
52
|
+
peak_mb = torch.cuda.max_memory_reserved() / (1024 * 1024)
|
|
53
|
+
print(f"SAM2 peak VRAM {peak_mb:.1f} MB")
|
|
54
|
+
if peak_mb >= 3000:
|
|
55
|
+
raise RuntimeError(f"SAM2 peak VRAM too high: {peak_mb:.1f} MB")
|
|
56
|
+
print("sam2 ok")
|
|
57
|
+
return 0
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
if __name__ == "__main__":
|
|
61
|
+
raise SystemExit(main())
|
|
@@ -0,0 +1,229 @@
|
|
|
1
|
+
#!/usr/bin/env python3
|
|
2
|
+
"""Evaluate Viterbi assignments against the Gate-0 truth-projected labels."""
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import argparse
|
|
6
|
+
import json
|
|
7
|
+
from pathlib import Path
|
|
8
|
+
from typing import Any
|
|
9
|
+
|
|
10
|
+
import numpy as np
|
|
11
|
+
|
|
12
|
+
from lr_viterbi import ARM_PAIRS, LEG_PAIRS, solve_lr_viterbi, state_permutation
|
|
13
|
+
|
|
14
|
+
COCO_TO_RIG = {
|
|
15
|
+
5: "LeftArm", 6: "RightArm", 7: "LeftForeArm", 8: "RightForeArm",
|
|
16
|
+
9: "LeftHand", 10: "RightHand", 11: "LeftUpLeg", 12: "RightUpLeg",
|
|
17
|
+
13: "LeftLeg", 14: "RightLeg", 15: "LeftFoot", 16: "RightFoot",
|
|
18
|
+
}
|
|
19
|
+
BODY = tuple(COCO_TO_RIG)
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def read_npz(path: Path) -> np.ndarray:
|
|
23
|
+
with np.load(path, allow_pickle=False) as data:
|
|
24
|
+
if "kp2d" not in data:
|
|
25
|
+
raise ValueError(f"{path}: missing kp2d")
|
|
26
|
+
return np.asarray(data["kp2d"], dtype=np.float64)
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
def _box(scene: dict[str, Any]) -> tuple[np.ndarray, np.ndarray, float]:
|
|
30
|
+
centre = np.asarray(scene.get("centre", scene.get("center")), dtype=np.float64)
|
|
31
|
+
half = np.asarray(scene.get("halfExtents"), dtype=np.float64)
|
|
32
|
+
if centre.shape != (3,) or half.shape != (3,):
|
|
33
|
+
if "min" not in scene or "max" not in scene:
|
|
34
|
+
raise ValueError("scene box needs centre/halfExtents or min/max")
|
|
35
|
+
lo, hi = np.asarray(scene["min"], dtype=np.float64), np.asarray(scene["max"], dtype=np.float64)
|
|
36
|
+
centre, half = (lo + hi) / 2, (hi - lo) / 2
|
|
37
|
+
yaw = float(scene.get("yawDeg", 0.0)) * np.pi / 180.0
|
|
38
|
+
if not np.isfinite(centre).all() or not np.isfinite(half).all() or (half <= 0).any():
|
|
39
|
+
raise ValueError("scene box has invalid dimensions")
|
|
40
|
+
return centre, half, yaw
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
def segment_hits_box(origin: np.ndarray, point: np.ndarray, scene: dict[str, Any]) -> bool:
|
|
44
|
+
centre, half, yaw = _box(scene)
|
|
45
|
+
c, s = np.cos(yaw), np.sin(yaw)
|
|
46
|
+
|
|
47
|
+
def local(p: np.ndarray) -> np.ndarray:
|
|
48
|
+
q = p - centre
|
|
49
|
+
return np.array([c * q[0] - s * q[2], q[1], s * q[0] + c * q[2]])
|
|
50
|
+
|
|
51
|
+
o, q = local(origin), local(point)
|
|
52
|
+
enter, exit = -np.inf, np.inf
|
|
53
|
+
for axis in range(3):
|
|
54
|
+
d = q[axis] - o[axis]
|
|
55
|
+
if abs(d) < 1e-12:
|
|
56
|
+
if abs(o[axis]) > half[axis]:
|
|
57
|
+
return False
|
|
58
|
+
continue
|
|
59
|
+
a, b = (-half[axis] - o[axis]) / d, (half[axis] - o[axis]) / d
|
|
60
|
+
enter, exit = max(enter, min(a, b)), min(exit, max(a, b))
|
|
61
|
+
return enter <= exit and exit > 1e-9 and enter < 1 - 1e-9
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
def truth_reference_error(study_item: Any, label: str) -> str | None:
|
|
65
|
+
"""Require todo-4's explicit truth-preferred arm/leg labels."""
|
|
66
|
+
if not isinstance(study_item, dict):
|
|
67
|
+
return f"{label}: missing study-2d item"
|
|
68
|
+
if study_item.get("status") != "ok":
|
|
69
|
+
return f"{label}: study-2d status is {study_item.get('status')!r}"
|
|
70
|
+
swaps = study_item.get("swaps")
|
|
71
|
+
frames = swaps.get("frames") if isinstance(swaps, dict) else None
|
|
72
|
+
if not isinstance(frames, dict) or not isinstance(frames.get("arms"), list) or not isinstance(frames.get("legs"), list):
|
|
73
|
+
return f"{label}: missing explicit swaps.frames.arms/legs truth labels"
|
|
74
|
+
if any(not isinstance(frame, int) or frame < 0 for group in ("arms", "legs") for frame in frames[group]):
|
|
75
|
+
return f"{label}: malformed truth swap frame labels"
|
|
76
|
+
return None
|
|
77
|
+
|
|
78
|
+
|
|
79
|
+
def expected_assignments(study_item: dict[str, Any], frames: int) -> np.ndarray:
|
|
80
|
+
"""Convert todo-4's per-group truth-preferred swap frames to permutations."""
|
|
81
|
+
arms = set(study_item.get("swaps", {}).get("frames", {}).get("arms", []))
|
|
82
|
+
legs = set(study_item.get("swaps", {}).get("frames", {}).get("legs", []))
|
|
83
|
+
out = np.empty((frames, 17), dtype=np.int64)
|
|
84
|
+
identity = state_permutation("identity")
|
|
85
|
+
arm_perm = state_permutation("arms_swap")
|
|
86
|
+
leg_perm = state_permutation("legs_swap")
|
|
87
|
+
full_perm = state_permutation("full_swap")
|
|
88
|
+
for t in range(frames):
|
|
89
|
+
if t in arms and t in legs:
|
|
90
|
+
out[t] = full_perm
|
|
91
|
+
elif t in arms:
|
|
92
|
+
out[t] = arm_perm
|
|
93
|
+
elif t in legs:
|
|
94
|
+
out[t] = leg_perm
|
|
95
|
+
else:
|
|
96
|
+
out[t] = identity
|
|
97
|
+
return out
|
|
98
|
+
|
|
99
|
+
|
|
100
|
+
def truth_visibility(item: dict[str, Any], frames: int, joints: dict[str, Any]) -> np.ndarray:
|
|
101
|
+
visible = np.ones((frames, 17), dtype=bool)
|
|
102
|
+
if not item.get("scene"):
|
|
103
|
+
return visible
|
|
104
|
+
scene = json.loads(Path(item["scene"]).read_text())
|
|
105
|
+
origin = np.asarray(joints.get("cameraPosition", [0, 0, 0]), dtype=np.float64)
|
|
106
|
+
if not np.isfinite(origin).all() or np.linalg.norm(origin) == 0:
|
|
107
|
+
camera_path = Path(item["dir"]) / item["variant"] / "camera.json"
|
|
108
|
+
camera = json.loads(camera_path.read_text())
|
|
109
|
+
pos = camera.get("position")
|
|
110
|
+
if isinstance(pos, dict):
|
|
111
|
+
origin = np.array([pos["x"], pos["y"], pos["z"]], dtype=np.float64)
|
|
112
|
+
else:
|
|
113
|
+
origin = np.asarray(pos, dtype=np.float64)
|
|
114
|
+
names = [joint["name"] if isinstance(joint, dict) else joint for joint in joints["joints"]]
|
|
115
|
+
index = {name: i for i, name in enumerate(names)}
|
|
116
|
+
world = np.asarray(joints["world"], dtype=np.float64)
|
|
117
|
+
for coco, rig_name in COCO_TO_RIG.items():
|
|
118
|
+
points = world[:frames, index[rig_name], :]
|
|
119
|
+
for t, point in enumerate(points):
|
|
120
|
+
visible[t, coco] = not segment_hits_box(origin, point, scene)
|
|
121
|
+
return visible
|
|
122
|
+
|
|
123
|
+
|
|
124
|
+
def evaluate_item(item: dict[str, Any], study_by_item: dict[str, Any], obs_root: Path, delta_px: float) -> dict[str, Any]:
|
|
125
|
+
label = f"{item['set']}/{item['name']}"
|
|
126
|
+
reference = study_by_item.get(label)
|
|
127
|
+
reference_error = truth_reference_error(reference, label)
|
|
128
|
+
if reference_error:
|
|
129
|
+
return {"item": label, "status": "missing-truth-reference", "error": reference_error}
|
|
130
|
+
obs_path = obs_root / item["set"] / item["name"] / "g5" / "obs.npz"
|
|
131
|
+
if not obs_path.exists():
|
|
132
|
+
return {"item": label, "status": "missing-obs", "obsPath": str(obs_path)}
|
|
133
|
+
variant_dir = Path(item["dir"]) / item["variant"]
|
|
134
|
+
joints = json.loads((variant_dir / "joints.json").read_text())
|
|
135
|
+
obs = read_npz(obs_path)
|
|
136
|
+
if obs.ndim != 3 or obs.shape[1] != 17 or obs.shape[2] < 2:
|
|
137
|
+
raise ValueError(f"{label}: kp2d has invalid shape {obs.shape}")
|
|
138
|
+
names = [joint["name"] if isinstance(joint, dict) else joint for joint in joints["joints"]]
|
|
139
|
+
index = {name: i for i, name in enumerate(names)}
|
|
140
|
+
uv = np.full((len(joints["world"]), 17, 2), np.nan, dtype=np.float64)
|
|
141
|
+
for coco, rig_name in COCO_TO_RIG.items():
|
|
142
|
+
uv[:, coco] = np.asarray(joints["uv"], dtype=np.float64)[: uv.shape[0], index[rig_name], :2]
|
|
143
|
+
frames = min(obs.shape[0], uv.shape[0])
|
|
144
|
+
obs, projected = obs[:frames], uv[:frames]
|
|
145
|
+
visible = truth_visibility(item, frames, joints)
|
|
146
|
+
result = solve_lr_viterbi(obs, projected, visibility=visible, deltaPx=delta_px)
|
|
147
|
+
expected = expected_assignments(reference, frames)
|
|
148
|
+
valid = np.isfinite(obs[..., :2]).all(axis=2) & (obs[..., 2] > 0 if obs.shape[2] >= 3 else True)
|
|
149
|
+
correct = raw_correct = total = 0
|
|
150
|
+
# Agreement is deliberately measured only on visible bilateral observations,
|
|
151
|
+
# never on box-occluded joints or absent detector points.
|
|
152
|
+
for t in range(frames):
|
|
153
|
+
for left, right in ARM_PAIRS + LEG_PAIRS:
|
|
154
|
+
if not (visible[t, left] and visible[t, right]):
|
|
155
|
+
continue
|
|
156
|
+
if not (valid[t, left] and valid[t, right]):
|
|
157
|
+
continue
|
|
158
|
+
for joint in (left, right):
|
|
159
|
+
total += 1
|
|
160
|
+
correct += int(result["assignments"][t, joint] == expected[t, joint])
|
|
161
|
+
raw_correct += int(joint == expected[t, joint])
|
|
162
|
+
return {
|
|
163
|
+
"item": label,
|
|
164
|
+
"status": "ok",
|
|
165
|
+
"frames": frames,
|
|
166
|
+
"deltaPx": delta_px,
|
|
167
|
+
"agreement": correct / total if total else None,
|
|
168
|
+
"correct": correct,
|
|
169
|
+
"rawAgreement": raw_correct / total if total else None,
|
|
170
|
+
"rawCorrect": raw_correct,
|
|
171
|
+
"visibleBilateralPairs": total // 2,
|
|
172
|
+
"visibleBilateralJointLabels": total,
|
|
173
|
+
"switches": result["switches"],
|
|
174
|
+
"stateCounts": result["state_counts"],
|
|
175
|
+
"ambiguousFrames": int(np.count_nonzero(result["ambiguous"])),
|
|
176
|
+
"medianMargin": float(np.median(result["margins"])) if frames else None,
|
|
177
|
+
}
|
|
178
|
+
|
|
179
|
+
|
|
180
|
+
def parse_args() -> argparse.Namespace:
|
|
181
|
+
here = Path(__file__).resolve()
|
|
182
|
+
root = next((parent for parent in here.parents if (parent / "evidence/obs/approved.json").exists()), Path.cwd())
|
|
183
|
+
parser = argparse.ArgumentParser(description=__doc__)
|
|
184
|
+
parser.add_argument("--summary", type=Path, default=root / "evidence/obs/study-2d/summary.json")
|
|
185
|
+
parser.add_argument("--approved", type=Path, default=root / "evidence/obs/approved.json")
|
|
186
|
+
parser.add_argument("--obs-root", type=Path, default=root / "evidence/obs/cache")
|
|
187
|
+
parser.add_argument("--delta-px", type=float, default=None)
|
|
188
|
+
parser.add_argument("--out", type=Path, default=None)
|
|
189
|
+
return parser.parse_args()
|
|
190
|
+
|
|
191
|
+
|
|
192
|
+
def main() -> int:
|
|
193
|
+
args = parse_args()
|
|
194
|
+
summary = json.loads(args.summary.read_text())
|
|
195
|
+
delta_px = args.delta_px if args.delta_px is not None else summary.get("deltaPx")
|
|
196
|
+
if not isinstance(delta_px, (int, float)) or not np.isfinite(delta_px) or delta_px <= 0:
|
|
197
|
+
raise SystemExit("eval_lr: summary must provide a finite positive deltaPx (or pass --delta-px)")
|
|
198
|
+
approved = json.loads(args.approved.read_text())
|
|
199
|
+
study_items = (summary.get("study2d") or {}).get("items", [])
|
|
200
|
+
study_by_item = {row.get("item"): row for row in study_items if isinstance(row, dict)}
|
|
201
|
+
truth_items = [item for item in approved["items"] if item.get("set") in {"gt", "cube"}]
|
|
202
|
+
results = [evaluate_item(item, study_by_item, args.obs_root, float(delta_px)) for item in truth_items]
|
|
203
|
+
valid = [row for row in results if row["status"] == "ok" and row["agreement"] is not None]
|
|
204
|
+
total = sum(row["visibleBilateralJointLabels"] for row in valid)
|
|
205
|
+
correct = sum(row["correct"] for row in valid)
|
|
206
|
+
raw_correct = sum(row["rawCorrect"] for row in valid)
|
|
207
|
+
output = {
|
|
208
|
+
"tool": "tools/track/py/eval_lr.py",
|
|
209
|
+
"summary": str(args.summary),
|
|
210
|
+
"deltaPx": float(delta_px),
|
|
211
|
+
"items": results,
|
|
212
|
+
"agreement": correct / total if total else None,
|
|
213
|
+
"correct": correct,
|
|
214
|
+
"rawAgreement": raw_correct / total if total else None,
|
|
215
|
+
"rawCorrect": raw_correct,
|
|
216
|
+
"visibleBilateralJointLabels": total,
|
|
217
|
+
"itemsOk": len(valid),
|
|
218
|
+
"itemsTotal": len(results),
|
|
219
|
+
"pass": bool(len(valid) == len(results) and total and correct / total >= 0.95),
|
|
220
|
+
}
|
|
221
|
+
out = args.out or args.summary.with_name("eval-lr.json")
|
|
222
|
+
out.parent.mkdir(parents=True, exist_ok=True)
|
|
223
|
+
out.write_text(json.dumps(output, indent=2) + "\n")
|
|
224
|
+
print(json.dumps(output, indent=2))
|
|
225
|
+
return 0 if output["pass"] else 1
|
|
226
|
+
|
|
227
|
+
|
|
228
|
+
if __name__ == "__main__":
|
|
229
|
+
raise SystemExit(main())
|
|
@@ -0,0 +1,294 @@
|
|
|
1
|
+
#!/usr/bin/env python3
|
|
2
|
+
"""Temporal left/right assignment for COCO-17 observations.
|
|
3
|
+
|
|
4
|
+
The detector's left/right labels are latent. This module intentionally knows
|
|
5
|
+
only geometry, confidence, and time; it has no appearance or colour inputs.
|
|
6
|
+
"""
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
from typing import Any, Iterable
|
|
10
|
+
|
|
11
|
+
import numpy as np
|
|
12
|
+
|
|
13
|
+
COCO_NAMES = (
|
|
14
|
+
"nose", "leftEye", "rightEye", "leftEar", "rightEar",
|
|
15
|
+
"leftShoulder", "rightShoulder", "leftElbow", "rightElbow",
|
|
16
|
+
"leftWrist", "rightWrist", "leftHip", "rightHip", "leftKnee",
|
|
17
|
+
"rightKnee", "leftAnkle", "rightAnkle",
|
|
18
|
+
)
|
|
19
|
+
|
|
20
|
+
# Pairs are COCO indices. Face pairs are changed only by the full-swap state.
|
|
21
|
+
ARM_PAIRS = ((5, 6), (7, 8), (9, 10))
|
|
22
|
+
LEG_PAIRS = ((11, 12), (13, 14), (15, 16))
|
|
23
|
+
FACE_PAIRS = ((1, 2), (3, 4))
|
|
24
|
+
BILATERAL_PAIRS = ARM_PAIRS + LEG_PAIRS
|
|
25
|
+
BODY_INDICES = tuple(i for pair in BILATERAL_PAIRS for i in pair)
|
|
26
|
+
|
|
27
|
+
STATE_NAMES = ("identity", "full_swap", "arms_swap", "legs_swap")
|
|
28
|
+
# The third bit is face-only. It is tied to both body swaps for full_swap.
|
|
29
|
+
STATE_BITS = {
|
|
30
|
+
"identity": (0, 0, 0),
|
|
31
|
+
"full_swap": (1, 1, 1),
|
|
32
|
+
"arms_swap": (1, 0, 0),
|
|
33
|
+
"legs_swap": (0, 1, 0),
|
|
34
|
+
}
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
def state_permutation(state: str | int) -> np.ndarray:
|
|
38
|
+
"""Return anatomical -> detector COCO indices for one permitted state."""
|
|
39
|
+
if isinstance(state, (int, np.integer)):
|
|
40
|
+
state = STATE_NAMES[int(state)]
|
|
41
|
+
if state not in STATE_BITS:
|
|
42
|
+
raise ValueError(f"unknown L/R state: {state!r}")
|
|
43
|
+
arm, leg, face = STATE_BITS[state]
|
|
44
|
+
out = np.arange(17, dtype=np.int64)
|
|
45
|
+
if arm:
|
|
46
|
+
for left, right in ARM_PAIRS:
|
|
47
|
+
out[left], out[right] = right, left
|
|
48
|
+
if leg:
|
|
49
|
+
for left, right in LEG_PAIRS:
|
|
50
|
+
out[left], out[right] = right, left
|
|
51
|
+
if face:
|
|
52
|
+
for left, right in FACE_PAIRS:
|
|
53
|
+
out[left], out[right] = right, left
|
|
54
|
+
return out
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
PERMUTATIONS = np.stack([state_permutation(name) for name in STATE_NAMES])
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
def huber(value: float | np.ndarray, delta: float) -> float | np.ndarray:
|
|
61
|
+
"""Huber rho, applied to a non-negative residual magnitude."""
|
|
62
|
+
if not np.isfinite(delta) or delta <= 0:
|
|
63
|
+
raise ValueError("deltaPx must be a finite positive number")
|
|
64
|
+
value = np.asarray(value, dtype=np.float64)
|
|
65
|
+
result = np.where(value <= delta, 0.5 * value * value, delta * (value - 0.5 * delta))
|
|
66
|
+
return float(result) if result.ndim == 0 else result
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
def _observations(kp2d: Any, projected: Any, visibility: Any | None) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]:
|
|
70
|
+
kp = np.asarray(kp2d, dtype=np.float64)
|
|
71
|
+
pred = np.asarray(projected, dtype=np.float64)
|
|
72
|
+
if kp.size == 0:
|
|
73
|
+
if pred.size == 0:
|
|
74
|
+
kp = np.empty((0, 17, 3), dtype=np.float64)
|
|
75
|
+
else:
|
|
76
|
+
pred = pred.reshape((-1, 17, 2))
|
|
77
|
+
kp = np.full((pred.shape[0], 17, 3), np.nan, dtype=np.float64)
|
|
78
|
+
elif kp.ndim == 2 and kp.shape == (17, 3):
|
|
79
|
+
kp = kp[None, ...]
|
|
80
|
+
elif kp.ndim == 1 and kp.size % (17 * 3) == 0:
|
|
81
|
+
kp = kp.reshape((-1, 17, 3))
|
|
82
|
+
if kp.ndim != 3 or kp.shape[1] != 17 or kp.shape[2] < 2:
|
|
83
|
+
raise ValueError("kp2d must have shape [frames, 17, 2+] (or be empty)")
|
|
84
|
+
|
|
85
|
+
if pred.size == 0 and kp.shape[0] == 0:
|
|
86
|
+
pred = np.empty((0, 17, 2), dtype=np.float64)
|
|
87
|
+
elif pred.ndim == 2 and pred.shape == (17, 2):
|
|
88
|
+
pred = pred[None, ...]
|
|
89
|
+
elif pred.ndim == 1 and pred.size % (17 * 2) == 0:
|
|
90
|
+
pred = pred.reshape((-1, 17, 2))
|
|
91
|
+
if pred.ndim != 3 or pred.shape[1:] != (17, 2):
|
|
92
|
+
raise ValueError("projected must have shape [frames, 17, 2]")
|
|
93
|
+
if kp.shape[0] != pred.shape[0]:
|
|
94
|
+
raise ValueError(f"kp2d/projected frame mismatch: {kp.shape[0]} vs {pred.shape[0]}")
|
|
95
|
+
|
|
96
|
+
conf = kp[..., 2] if kp.shape[2] >= 3 else np.ones(kp.shape[:2], dtype=np.float64)
|
|
97
|
+
conf = np.where(np.isfinite(conf), np.clip(conf, 0.0, 1.0), 0.0)
|
|
98
|
+
detector_weight = np.where(np.isfinite(kp[..., :2]).all(axis=-1), conf, 0.0)
|
|
99
|
+
anatomical_weight = np.isfinite(pred).all(axis=-1).astype(np.float64)
|
|
100
|
+
if visibility is not None:
|
|
101
|
+
vis = np.asarray(visibility, dtype=bool)
|
|
102
|
+
if vis.shape != anatomical_weight.shape:
|
|
103
|
+
raise ValueError(f"visibility must have shape {anatomical_weight.shape}, got {vis.shape}")
|
|
104
|
+
anatomical_weight *= vis
|
|
105
|
+
# Visibility belongs to the rendered anatomical joint i. Detector-side
|
|
106
|
+
# confidence/validity remains indexed by h(i) and is permuted later.
|
|
107
|
+
return kp, pred, detector_weight, anatomical_weight
|
|
108
|
+
|
|
109
|
+
|
|
110
|
+
def _frame_scales(t: int, torso_scale: Any | None) -> np.ndarray:
|
|
111
|
+
if torso_scale is None:
|
|
112
|
+
return np.ones(t, dtype=np.float64)
|
|
113
|
+
scale = np.asarray(torso_scale, dtype=np.float64)
|
|
114
|
+
if scale.ndim == 0:
|
|
115
|
+
scale = np.full(t, float(scale), dtype=np.float64)
|
|
116
|
+
elif scale.ndim == 1 and scale.shape == (t,):
|
|
117
|
+
scale = scale.copy()
|
|
118
|
+
else:
|
|
119
|
+
raise ValueError("torso_scale must be a scalar or one value per frame")
|
|
120
|
+
if not np.isfinite(scale).all() or (scale <= 0).any():
|
|
121
|
+
raise ValueError("torso_scale must contain finite positive values")
|
|
122
|
+
return scale
|
|
123
|
+
|
|
124
|
+
|
|
125
|
+
def _emissions(kp: np.ndarray, pred: np.ndarray, detector_weight: np.ndarray, anatomical_weight: np.ndarray, scales: np.ndarray, delta: float, complexity_weight: float) -> tuple[np.ndarray, np.ndarray]:
|
|
126
|
+
t = kp.shape[0]
|
|
127
|
+
emissions = np.zeros((t, len(STATE_NAMES)), dtype=np.float64)
|
|
128
|
+
identity_terms: list[np.ndarray] = []
|
|
129
|
+
for state_index, permutation in enumerate(PERMUTATIONS):
|
|
130
|
+
observed = kp[:, permutation, :2]
|
|
131
|
+
residual = np.linalg.norm(pred - observed, axis=2) / scales[:, None]
|
|
132
|
+
residual = np.nan_to_num(residual, nan=0.0, posinf=0.0, neginf=0.0)
|
|
133
|
+
robust = np.asarray(huber(residual, delta))
|
|
134
|
+
joint_weight = anatomical_weight * detector_weight[:, permutation]
|
|
135
|
+
weighted = robust * joint_weight
|
|
136
|
+
emissions[:, state_index] = weighted.sum(axis=1) + complexity_weight * sum(STATE_BITS[STATE_NAMES[state_index]])
|
|
137
|
+
if state_index == 0:
|
|
138
|
+
identity_terms.append(robust[joint_weight > 0])
|
|
139
|
+
return emissions, np.concatenate(identity_terms) if identity_terms else np.empty(0, dtype=np.float64)
|
|
140
|
+
|
|
141
|
+
|
|
142
|
+
def _transitions(kp: np.ndarray, pred: np.ndarray, detector_weight: np.ndarray, anatomical_weight: np.ndarray, scales: np.ndarray, delta: float, switch_penalty: float, lambda_cont: float) -> np.ndarray:
|
|
143
|
+
"""Return [frame-1, previous_state, current_state] transition costs."""
|
|
144
|
+
t = kp.shape[0]
|
|
145
|
+
out = np.zeros((max(0, t - 1), len(STATE_NAMES), len(STATE_NAMES)), dtype=np.float64)
|
|
146
|
+
predicted_delta = pred[1:] - pred[:-1]
|
|
147
|
+
for previous, prev_perm in enumerate(PERMUTATIONS):
|
|
148
|
+
for current, current_perm in enumerate(PERMUTATIONS):
|
|
149
|
+
observed_delta = kp[1:, current_perm, :2] - kp[:-1, prev_perm, :2]
|
|
150
|
+
residual = np.linalg.norm(observed_delta - predicted_delta, axis=2) / scales[1:, None]
|
|
151
|
+
residual = np.nan_to_num(residual, nan=0.0, posinf=0.0, neginf=0.0)
|
|
152
|
+
robust = np.asarray(huber(residual, delta))
|
|
153
|
+
detector_pair_weights = np.minimum(detector_weight[1:, current_perm], detector_weight[:-1, prev_perm])
|
|
154
|
+
anatomical_pair_weights = np.minimum(anatomical_weight[1:], anatomical_weight[:-1])
|
|
155
|
+
continuous = (robust * detector_pair_weights * anatomical_pair_weights).sum(axis=1)
|
|
156
|
+
changed = sum(a != b for a, b in zip(STATE_BITS[STATE_NAMES[previous]], STATE_BITS[STATE_NAMES[current]]))
|
|
157
|
+
out[:, previous, current] = switch_penalty * changed + lambda_cont * continuous
|
|
158
|
+
return out
|
|
159
|
+
|
|
160
|
+
|
|
161
|
+
def solve_lr_viterbi(
|
|
162
|
+
kp2d: Any,
|
|
163
|
+
projected: Any,
|
|
164
|
+
*,
|
|
165
|
+
visibility: Any | None = None,
|
|
166
|
+
torso_scale: Any | None = None,
|
|
167
|
+
deltaPx: float = 1.0,
|
|
168
|
+
switch_penalty: float | None = None,
|
|
169
|
+
lambda_cont: float = 1.0,
|
|
170
|
+
hysteresis: float = 0.1,
|
|
171
|
+
complexity_weight: float | None = None,
|
|
172
|
+
) -> dict[str, Any]:
|
|
173
|
+
"""Solve the permitted COCO left/right assignment states over a clip.
|
|
174
|
+
|
|
175
|
+
``path[t]`` is an integer in ``STATE_NAMES`` order. ``assignment[t, i]``
|
|
176
|
+
is the detector index to use for anatomical joint ``i``. Invalid/NaN
|
|
177
|
+
observations simply have zero weight; this makes an all-zero-confidence
|
|
178
|
+
clip deterministic (identity, all ambiguous) rather than exceptional.
|
|
179
|
+
"""
|
|
180
|
+
if not np.isfinite(deltaPx) or deltaPx <= 0:
|
|
181
|
+
raise ValueError("deltaPx must be a finite positive number")
|
|
182
|
+
if lambda_cont < 0 or hysteresis < 0:
|
|
183
|
+
raise ValueError("lambda_cont and hysteresis must be non-negative")
|
|
184
|
+
kp, pred, detector_weight, anatomical_weight = _observations(kp2d, projected, visibility)
|
|
185
|
+
t = kp.shape[0]
|
|
186
|
+
if t == 0:
|
|
187
|
+
return {
|
|
188
|
+
"path": np.empty(0, dtype=np.int64), "state_names": STATE_NAMES,
|
|
189
|
+
"assignments": np.empty((0, 17), dtype=np.int64),
|
|
190
|
+
"margins": np.empty(0, dtype=np.float64), "ambiguous": np.empty(0, dtype=bool),
|
|
191
|
+
"emission": np.empty((0, 4), dtype=np.float64), "switch_penalty": 0.0,
|
|
192
|
+
"state_counts": {name: 0 for name in STATE_NAMES}, "switches": 0,
|
|
193
|
+
}
|
|
194
|
+
|
|
195
|
+
scales = _frame_scales(t, torso_scale)
|
|
196
|
+
# The plan's switch penalty is four times the median visible emission. Use
|
|
197
|
+
# the robust per-joint identity terms, not an aggregate frame cost, so a
|
|
198
|
+
# short swap run cannot inflate its own penalty.
|
|
199
|
+
_, identity_terms = _emissions(kp, pred, detector_weight, anatomical_weight, scales, float(deltaPx), 0.0)
|
|
200
|
+
median_emission = float(np.median(identity_terms)) if identity_terms.size else 0.0
|
|
201
|
+
penalty = float(4.0 * median_emission if switch_penalty is None else switch_penalty)
|
|
202
|
+
if not np.isfinite(penalty) or penalty < 0:
|
|
203
|
+
raise ValueError("switch_penalty must be finite and non-negative")
|
|
204
|
+
if complexity_weight is None:
|
|
205
|
+
complexity_weight = 0.01 * median_emission
|
|
206
|
+
emissions, _ = _emissions(kp, pred, detector_weight, anatomical_weight, scales, float(deltaPx), float(complexity_weight))
|
|
207
|
+
transitions = _transitions(kp, pred, detector_weight, anatomical_weight, scales, float(deltaPx), penalty, float(lambda_cont))
|
|
208
|
+
|
|
209
|
+
states = len(STATE_NAMES)
|
|
210
|
+
forward = np.full((t, states), np.inf, dtype=np.float64)
|
|
211
|
+
backpointer = np.zeros((t, states), dtype=np.int64)
|
|
212
|
+
forward[0] = emissions[0]
|
|
213
|
+
for frame in range(1, t):
|
|
214
|
+
for current in range(states):
|
|
215
|
+
values = forward[frame - 1] + transitions[frame - 1, :, current]
|
|
216
|
+
previous = int(np.argmin(values))
|
|
217
|
+
backpointer[frame, current] = previous
|
|
218
|
+
forward[frame, current] = emissions[frame, current] + values[previous]
|
|
219
|
+
raw_path = np.empty(t, dtype=np.int64)
|
|
220
|
+
raw_path[-1] = int(np.argmin(forward[-1]))
|
|
221
|
+
for frame in range(t - 1, 0, -1):
|
|
222
|
+
raw_path[frame - 1] = backpointer[frame, raw_path[frame]]
|
|
223
|
+
|
|
224
|
+
backward = np.zeros((t, states), dtype=np.float64)
|
|
225
|
+
for frame in range(t - 2, -1, -1):
|
|
226
|
+
for previous in range(states):
|
|
227
|
+
backward[frame, previous] = np.min(transitions[frame, previous] + emissions[frame + 1] + backward[frame + 1])
|
|
228
|
+
conditioned = forward + backward - np.min(forward[-1])
|
|
229
|
+
|
|
230
|
+
def assignment_margin(frame: int, chosen: int) -> float:
|
|
231
|
+
"""Margin against the best state that changes an active assignment."""
|
|
232
|
+
chosen_perm = PERMUTATIONS[chosen]
|
|
233
|
+
anatomical_active = anatomical_weight[frame] > 0
|
|
234
|
+
best = np.inf
|
|
235
|
+
for alternative, alternative_perm in enumerate(PERMUTATIONS):
|
|
236
|
+
if alternative == chosen:
|
|
237
|
+
continue
|
|
238
|
+
affected = anatomical_active & (alternative_perm != chosen_perm)
|
|
239
|
+
if not np.any(affected):
|
|
240
|
+
# Ties in inactive groups are not assignment alternatives.
|
|
241
|
+
continue
|
|
242
|
+
detector_active = (detector_weight[frame, chosen_perm] > 0) | (detector_weight[frame, alternative_perm] > 0)
|
|
243
|
+
affected &= detector_active
|
|
244
|
+
count = int(np.count_nonzero(affected))
|
|
245
|
+
if not count:
|
|
246
|
+
continue
|
|
247
|
+
best = min(best, float((conditioned[frame, alternative] - conditioned[frame, chosen]) / count))
|
|
248
|
+
return max(0.0, best) if np.isfinite(best) else 0.0
|
|
249
|
+
|
|
250
|
+
# Hysteresis compares the proposed assignment with the assignment it would
|
|
251
|
+
# retain, not with an unrelated state that ties on inactive groups.
|
|
252
|
+
path = raw_path.copy()
|
|
253
|
+
for frame in range(1, t):
|
|
254
|
+
proposed, retained = int(raw_path[frame]), int(path[frame - 1])
|
|
255
|
+
if proposed == retained:
|
|
256
|
+
continue
|
|
257
|
+
proposed_perm, retained_perm = PERMUTATIONS[proposed], PERMUTATIONS[retained]
|
|
258
|
+
active = anatomical_weight[frame] > 0
|
|
259
|
+
active &= (detector_weight[frame, proposed_perm] > 0) | (detector_weight[frame, retained_perm] > 0)
|
|
260
|
+
affected = active & (proposed_perm != retained_perm)
|
|
261
|
+
count = int(np.count_nonzero(affected))
|
|
262
|
+
if not count:
|
|
263
|
+
path[frame] = retained
|
|
264
|
+
continue
|
|
265
|
+
improvement_per_joint = float((conditioned[frame, retained] - conditioned[frame, proposed]) / count)
|
|
266
|
+
if improvement_per_joint < hysteresis:
|
|
267
|
+
path[frame] = retained
|
|
268
|
+
|
|
269
|
+
margins = np.asarray([assignment_margin(frame, int(path[frame])) for frame in range(t)], dtype=np.float64)
|
|
270
|
+
ambiguous = margins < 0.05
|
|
271
|
+
assignments = PERMUTATIONS[path]
|
|
272
|
+
counts = {name: int(np.count_nonzero(path == index)) for index, name in enumerate(STATE_NAMES)}
|
|
273
|
+
return {
|
|
274
|
+
"path": path,
|
|
275
|
+
"raw_path": raw_path,
|
|
276
|
+
"state_names": STATE_NAMES,
|
|
277
|
+
"assignments": assignments,
|
|
278
|
+
"margins": margins,
|
|
279
|
+
"ambiguous": ambiguous,
|
|
280
|
+
"emission": emissions,
|
|
281
|
+
"switch_penalty": penalty,
|
|
282
|
+
"median_visible_emission": median_emission,
|
|
283
|
+
"state_counts": counts,
|
|
284
|
+
"switches": int(np.count_nonzero(path[1:] != path[:-1])),
|
|
285
|
+
}
|
|
286
|
+
|
|
287
|
+
|
|
288
|
+
# Short aliases make the pure function convenient for callers and tests.
|
|
289
|
+
viterbi_lr = solve_lr_viterbi
|
|
290
|
+
solve_viterbi = solve_lr_viterbi
|
|
291
|
+
|
|
292
|
+
|
|
293
|
+
if __name__ == "__main__":
|
|
294
|
+
raise SystemExit("lr_viterbi.py is a library; use eval_lr.py for file evaluation")
|