wapr 0.0.1__py3-none-any.whl

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.
Files changed (83) hide show
  1. wapr/__init__.py +18 -0
  2. wapr/bootstrap.py +311 -0
  3. wapr/bop.py +217 -0
  4. wapr/det2d.py +2833 -0
  5. wapr/det2d_assets.json +40 -0
  6. wapr/download_assets.py +649 -0
  7. wapr/download_route.py +271 -0
  8. wapr/estimator.py +1386 -0
  9. wapr/export_batch_engine.py +133 -0
  10. wapr/export_engines.py +510 -0
  11. wapr/export_tracking_engine.py +123 -0
  12. wapr/fonts/Apache-2.0.txt +202 -0
  13. wapr/fonts/DejaVuSans-Bold.ttf +0 -0
  14. wapr/fonts/DejaVuSans.ttf +0 -0
  15. wapr/fonts/LICENSE_DEJAVU +99 -0
  16. wapr/fonts/wqy-microhei-copyright.txt +91 -0
  17. wapr/fonts/wqy-microhei.ttc +0 -0
  18. wapr/frame.py +421 -0
  19. wapr/installation.py +437 -0
  20. wapr/mesh_extent.py +197 -0
  21. wapr/model_metadata.py +107 -0
  22. wapr/nets.py +2082 -0
  23. wapr/ogl.py +3219 -0
  24. wapr/ogl_native/CMakeLists.txt +82 -0
  25. wapr/ogl_native/cpp/asset_manager.cpp +528 -0
  26. wapr/ogl_native/cpp/asset_manager.h +77 -0
  27. wapr/ogl_native/cpp/batch_renderer.cpp +1473 -0
  28. wapr/ogl_native/cpp/batch_renderer.h +153 -0
  29. wapr/ogl_native/cpp/egl_context.cpp +174 -0
  30. wapr/ogl_native/cpp/egl_context.h +42 -0
  31. wapr/ogl_native/cpp/gl_loader.cpp +103 -0
  32. wapr/ogl_native/cpp/gl_loader.h +180 -0
  33. wapr/ogl_native/cpp/gpu_render_runtime.cpp +696 -0
  34. wapr/ogl_native/cpp/gpu_render_runtime.h +128 -0
  35. wapr/ogl_native/cpp/projection.cpp +206 -0
  36. wapr/ogl_native/cpp/projection.h +89 -0
  37. wapr/ogl_native/cpp/shader_utils.cpp +75 -0
  38. wapr/ogl_native/cpp/shader_utils.h +25 -0
  39. wapr/ogl_native/cpp/types.h +284 -0
  40. wapr/ogl_native/cuda/cuda_device_guard.h +47 -0
  41. wapr/ogl_native/cuda/cuda_distort_context.cpp +124 -0
  42. wapr/ogl_native/cuda/cuda_distort_context.h +47 -0
  43. wapr/ogl_native/cuda/cuda_pack_context.cpp +586 -0
  44. wapr/ogl_native/cuda/cuda_pack_context.h +121 -0
  45. wapr/ogl_native/cuda/distort_remap.cu +205 -0
  46. wapr/ogl_native/cuda/distort_remap.cuh +29 -0
  47. wapr/ogl_native/cuda/fill_tile_instances.cu +212 -0
  48. wapr/ogl_native/cuda/fill_tile_instances.h +41 -0
  49. wapr/ogl_native/cuda/pack_outputs.cu +169 -0
  50. wapr/ogl_native/cuda/pack_outputs.cuh +48 -0
  51. wapr/ogl_native/python/gpu_render_pybind.cpp +577 -0
  52. wapr/ogl_native/shaders/lit_mrt.frag +219 -0
  53. wapr/ogl_native/shaders/lit_mrt.vert +155 -0
  54. wapr/pose_groups.py +24 -0
  55. wapr/raster_fallback.py +707 -0
  56. wapr/recipe.py +190 -0
  57. wapr/reconstruction_setup.py +260 -0
  58. wapr/region_tracking.py +193 -0
  59. wapr/resources.py +58 -0
  60. wapr/runtime_assets/groundingdino/LICENSE +201 -0
  61. wapr/runtime_assets/groundingdino/groundingdino_swinb.json +49 -0
  62. wapr/runtime_assets/groundingdino/groundingdino_swint.json +45 -0
  63. wapr/runtime_assets/groundingdino/ms_deform_attn.py +340 -0
  64. wapr/runtime_assets/patches/dinov2_python38_annotations.patch +90 -0
  65. wapr/runtime_assets/patches/sam3d_wapr_compat.patch +195 -0
  66. wapr/scene_files.py +296 -0
  67. wapr/source_setup.py +328 -0
  68. wapr/suppression.py +173 -0
  69. wapr/tracking.py +268 -0
  70. wapr/view.py +1101 -0
  71. wapr/weight_license/CC-BY-ND-4.0.txt +393 -0
  72. wapr/weight_license/WEIGHTS_LICENSE.txt +50 -0
  73. wapr-0.0.1.dist-info/METADATA +233 -0
  74. wapr-0.0.1.dist-info/RECORD +83 -0
  75. wapr-0.0.1.dist-info/WHEEL +5 -0
  76. wapr-0.0.1.dist-info/licenses/AUTHORS.md +42 -0
  77. wapr-0.0.1.dist-info/licenses/LICENSE +504 -0
  78. wapr-0.0.1.dist-info/licenses/THIRD_PARTY_NOTICES.txt +347 -0
  79. wapr-0.0.1.dist-info/licenses/WEIGHTS_LICENSE.txt +50 -0
  80. wapr-0.0.1.dist-info/licenses/wapr/fonts/Apache-2.0.txt +202 -0
  81. wapr-0.0.1.dist-info/licenses/wapr/fonts/LICENSE_DEJAVU +99 -0
  82. wapr-0.0.1.dist-info/licenses/wapr/fonts/wqy-microhei-copyright.txt +91 -0
  83. wapr-0.0.1.dist-info/top_level.txt +1 -0
wapr/__init__.py ADDED
@@ -0,0 +1,18 @@
1
+ # Author: Yulin Wang (yulinwang@seu.edu.cn)
2
+ # School of Mechanical Engineering, Southeast University, China
3
+ # Copyright (c) 2026 Yulin Wang. All rights reserved, except as granted under LICENSE.
4
+ # SPDX-License-Identifier: LGPL-2.1-only
5
+ # 作者与版权人:Yulin Wang;使用、修改与再分发须遵守项目 LICENSE。
6
+ # Third-party portions retain their original notices and terms; see THIRD_PARTY_NOTICES.txt.
7
+
8
+ # Public entry: from wapr import WAPREstimator.
9
+ # 对外入口:from wapr import WAPREstimator。
10
+ __all__ = ["WAPREstimator"]
11
+
12
+
13
+ def __getattr__(name):
14
+ """Load inference after dependency preparation. / 依赖准备后再载入推理模块。"""
15
+ if name == "WAPREstimator":
16
+ from wapr.estimator import WAPREstimator
17
+ return WAPREstimator
18
+ raise AttributeError(name)
wapr/bootstrap.py ADDED
@@ -0,0 +1,311 @@
1
+ # Author: Yulin Wang (yulinwang@seu.edu.cn)
2
+ # SPDX-License-Identifier: LGPL-2.1-only
3
+ """Prepare an installed WAPR source wheel in the existing Python environment.
4
+
5
+ 在当前 Python 环境准备 WAPR 源码 wheel;不创建环境,更换已有包须用户明确同意。
6
+ Run / 运行: python -m wapr.bootstrap
7
+ """
8
+ import ctypes.util
9
+ import argparse
10
+ import importlib.util
11
+ import json
12
+ import os
13
+ import re
14
+ import shutil
15
+ import subprocess
16
+ import sys
17
+ import tempfile
18
+ from pathlib import Path
19
+ from urllib.error import HTTPError, URLError
20
+ from urllib.request import ProxyHandler, Request, build_opener
21
+
22
+
23
+ _prepared_optional = set()
24
+
25
+
26
+ def check_runtime():
27
+ """Inspect the existing core environment without installing anything.
28
+
29
+ 只检查当前核心环境,不安装包、不编译、不下载权重。
30
+ """
31
+ from importlib import metadata
32
+ versions = {}
33
+ for name in ("torch", "torchvision", "torchaudio", "numpy", "trimesh", "kornia", "onnx"):
34
+ try:
35
+ versions[name] = metadata.version(name)
36
+ except metadata.PackageNotFoundError:
37
+ versions[name] = None
38
+ result = {"feature": "core", "python": sys.version.split()[0], "executable": sys.executable,
39
+ "platform": sys.platform, "versions": versions, "status": "needs_preparation"}
40
+ try:
41
+ import torch
42
+ result.update(torch_cuda=torch.version.cuda, cuda_available=torch.cuda.is_available())
43
+ if torch.cuda.is_available():
44
+ result.update(gpu=torch.cuda.get_device_name(0), gpu_memory_gb=round(
45
+ torch.cuda.get_device_properties(0).total_memory / 1024 ** 3, 2))
46
+ result["nvcc"] = shutil.which("nvcc") or (
47
+ "/usr/local/cuda/bin/nvcc" if os.path.isfile("/usr/local/cuda/bin/nvcc") else None)
48
+ if sys.platform != "linux" or not torch.cuda.is_available():
49
+ result.update(status="blocked", reason="OGL requires Linux and working CUDA / OGL 需要 Linux 与可用 CUDA")
50
+ elif all(versions[name] is not None for name in ("numpy", "trimesh", "kornia")) and result["nvcc"]:
51
+ result["status"] = "dependencies_present"
52
+ result["note"] = "Presence is not inference validation / 依赖存在不等于推理验证通过"
53
+ except (ImportError, OSError) as error:
54
+ result.update(status="blocked", reason="Existing PyTorch is unavailable / 当前 PyTorch 不可用: " + type(error).__name__)
55
+ return result
56
+
57
+
58
+ def prepare_feature(feature, allow_replacement=None, check_only=False):
59
+ """Prepare one requested feature, preserving unrelated optional stacks.
60
+
61
+ 只准备被请求的功能;更换已有库须明确同意,不安装无关可选功能。
62
+ """
63
+ from wapr.installation import install_requirements, prepare_optional
64
+ if feature == "core":
65
+ if check_only:
66
+ return check_runtime()
67
+ prepare_runtime(allow_replacement=allow_replacement)
68
+ return {"feature": "core", "status": "ready"}
69
+ if feature == "sam3d":
70
+ from wapr.reconstruction_setup import prepare_reconstruction
71
+ return prepare_reconstruction(allow_replacement=allow_replacement, check_only=check_only)
72
+ if feature == "compatible":
73
+ # Resolve first; replacing an existing package still requires approval.
74
+ # 先逐项解析;更换已有包仍须明确同意,无法准备的功能单独报告。
75
+ results = []
76
+ for name in ("dinov2", "det2d", "robot", "sam2", "roma", "qwen", "sam3d", "unipose9d"):
77
+ plan = prepare_feature(name, allow_replacement=False, check_only=True)
78
+ if not check_only and plan.get("status") in ("ready", "installable", "dependencies_present", "approval_required", "needs_source"):
79
+ try:
80
+ plan = prepare_feature(name, allow_replacement=allow_replacement)
81
+ except (RuntimeError, OSError, subprocess.SubprocessError) as error:
82
+ plan = {"feature": name, "status": "blocked", "reason": str(error)}
83
+ results.append(plan)
84
+ if plan.get("status") == "restart_required":
85
+ return {"feature": feature, "status": "restart_required", "features": results,
86
+ "reason": "Restart Python before preparing further features / 请重启 Python 后继续准备其他功能"}
87
+ incomplete = any(item.get("status") not in ("ready", "installed", "prepared") for item in results)
88
+ return {"feature": feature, "status": "checked" if check_only else ("partial" if incomplete else "prepared"), "features": results}
89
+ if not check_only and feature in _prepared_optional:
90
+ return {"feature": feature, "status": "ready"}
91
+ source = None
92
+ if feature in ("sam2", "roma"):
93
+ from wapr.source_setup import prepare_source
94
+ if check_only:
95
+ from wapr.resources import resource_root, source_checkout
96
+ source_name = "sam2" if feature == "sam2" else "RoMa"
97
+ source_parent = os.path.join(resource_root(), "third_party" if source_checkout else "sources")
98
+ candidate = os.path.join(source_parent, source_name)
99
+ if os.path.isdir(candidate):
100
+ source = prepare_source(feature)
101
+ else:
102
+ source = prepare_source(feature)
103
+ result = prepare_optional(feature, allow_replacement=allow_replacement, check_only=check_only,
104
+ source_requirement=source)
105
+ if check_only:
106
+ if result.get("status") == "ready" and result.get("install"):
107
+ result["status"] = "installable"
108
+ return result
109
+ if result.get("status") not in ("ready", "installed"):
110
+ return result
111
+ if feature == "unipose9d" and result.get("existing_source_api"):
112
+ _prepared_optional.add(feature)
113
+ return result
114
+ # Source-based projects stay outside site-packages and outside the wheel.
115
+ # 源码项目放在资源目录,不打进 wheel;使用当前解释器解析其实际依赖。
116
+ if feature in ("dinov2", "det2d", "sam2", "roma", "unipose9d"):
117
+ from wapr.source_setup import prepare_source
118
+ if source is None:
119
+ source = prepare_source(feature)
120
+ result["source"] = source
121
+ if feature == "det2d":
122
+ parent = os.path.dirname(source)
123
+ source_paths = [source, os.path.join(parent, "ultralytics"), os.path.join(parent, "dinov2")]
124
+ elif feature == "unipose9d":
125
+ source_paths = [os.path.join(source, "infer")]
126
+ else:
127
+ source_paths = [source]
128
+ for source_path in source_paths:
129
+ if source_path not in sys.path:
130
+ sys.path.insert(0, source_path)
131
+ if feature == "unipose9d":
132
+ import importlib
133
+ module = importlib.import_module("unipose9d_inference")
134
+ for name in ("estimate_pose", "load_pose_model", "set_seed"):
135
+ if not callable(getattr(module, name, None)):
136
+ raise RuntimeError("UniPose9D API missing / UniPose9D API 缺失: " + name)
137
+ # This marker means installation finished, not that model inference was tested.
138
+ # 此标记只表示安装阶段完成,不表示模型推理已经验证。
139
+ _prepared_optional.add(feature)
140
+ result["status"] = "ready"
141
+ return result
142
+
143
+
144
+ def ensure_optional(feature):
145
+ """First-call preparation; decline or conflict stops this feature.
146
+
147
+ 首次调用时准备;拒绝更换或依赖冲突时退出该功能。
148
+ """
149
+ result = prepare_feature(feature)
150
+ if result.get("status") != "ready":
151
+ raise RuntimeError("Optional feature preparation stopped / 可选功能准备已停止: " + json.dumps(result, ensure_ascii=False))
152
+ return result
153
+
154
+
155
+ def prepare_runtime(allow_replacement=None):
156
+ """Install missing pose dependencies and compile for the existing CUDA.
157
+
158
+ 安装缺失的位姿依赖,按当前 CUDA 编译;默认保留已有库,更换须明确同意。
159
+ """
160
+ if sys.platform != "linux":
161
+ raise RuntimeError("This native renderer requires Linux / 本地渲染器需要 Linux")
162
+ import torch
163
+ from importlib import metadata
164
+ if not torch.cuda.is_available():
165
+ raise RuntimeError("Existing PyTorch cannot use CUDA / 已有 PyTorch 无法使用 CUDA")
166
+ original_torch = torch.__version__
167
+ original_cuda = torch.version.cuda
168
+ missing = []
169
+ # OpenCV 4.12+ requires NumPy 2; retain preinstalled NumPy 1 environments.
170
+ # OpenCV 4.12 之后要求 NumPy 2;保留镜像自带 NumPy 1 时选择兼容版本范围。
171
+ opencv_package = "opencv-python-headless"
172
+ try:
173
+ if int(metadata.version("numpy").split(".")[0]) < 2:
174
+ opencv_package += "<4.12"
175
+ except metadata.PackageNotFoundError:
176
+ pass
177
+ for module, package in [("numpy", "numpy"), ("scipy", "scipy"), ("trimesh", "trimesh"), ("PIL", "Pillow"),
178
+ ("huggingface_hub", "huggingface-hub"), ("kornia", "kornia"),
179
+ ("cv2", opencv_package), ("pybind11", "pybind11>=2.10")]:
180
+ if importlib.util.find_spec(module) is None:
181
+ missing.append(package)
182
+ # TensorRT provides separate CUDA families. Select from the existing torch
183
+ # CUDA runtime; its own packages are downloaded rather than bundled here.
184
+ # TensorRT 分 CUDA 家族发行;按已有 torch 的 CUDA 运行时选包,不打进本 wheel。
185
+ cuda_major = int(original_cuda.split(".")[0])
186
+ if cuda_major not in (11, 12, 13):
187
+ raise RuntimeError("Unverified TensorRT CUDA family / 未验证的 TensorRT CUDA 家族")
188
+ if importlib.util.find_spec("tensorrt") is None:
189
+ missing.append("tensorrt-cu%d>=10,<11" % cuda_major)
190
+ if importlib.util.find_spec("onnx") is None:
191
+ missing.append("onnx")
192
+ cmake_path = shutil.which("cmake")
193
+ cmake_version = (0, 0)
194
+ if cmake_path is not None:
195
+ cmake_output = subprocess.check_output([cmake_path, "--version"], text=True)
196
+ cmake_match = re.search(r"version (\d+)\.(\d+)", cmake_output)
197
+ if cmake_match is not None:
198
+ cmake_version = tuple(int(value) for value in cmake_match.groups())
199
+ if cmake_version < (3, 18):
200
+ missing.append("cmake>=3.18")
201
+ print("WAPR_TARGET", {"python": sys.version.split()[0], "executable": sys.executable,
202
+ "torch": original_torch, "torch_cuda": original_cuda,
203
+ "gpu": torch.cuda.get_device_name(0), "missing": missing}, flush=True)
204
+ if missing:
205
+ from wapr.installation import install_requirements
206
+ dependency_result = install_requirements(missing, allow_replacement=allow_replacement)
207
+ if dependency_result.get("status") not in ("ready", "installed"):
208
+ raise RuntimeError("Core preparation stopped / 核心环境准备已停止: " + json.dumps(dependency_result, ensure_ascii=False))
209
+ # Missing EGL/GL development headers are system prerequisites, not wheel contents.
210
+ # EGL/GL 开发头文件是系统前置条件,不打进 wheel;缺项才安装。
211
+ system_packages = []
212
+ for header, package in [("/usr/include/EGL/egl.h", "libegl1-mesa-dev"),
213
+ ("/usr/include/GL/gl.h", "libgl1-mesa-dev")]:
214
+ if not os.path.isfile(header):
215
+ system_packages.append(package)
216
+ if shutil.which("g++") is None:
217
+ system_packages.append("g++")
218
+ if system_packages:
219
+ if os.geteuid() != 0 or shutil.which("apt-get") is None:
220
+ raise RuntimeError("Install system prerequisites / 请安装系统前置依赖: " + " ".join(system_packages))
221
+ subprocess.check_call(["apt-get", "update"])
222
+ subprocess.check_call(["apt-get", "install", "-y"] + system_packages)
223
+ if metadata.version("torch") != original_torch:
224
+ raise RuntimeError("Restart Python after an approved PyTorch replacement / 同意更换 PyTorch 后,请重启 Python 再准备运行环境")
225
+ from wapr.ogl import ensure_ogl
226
+ ensure_ogl()
227
+ print("WAPR_RUNTIME_READY", {"torch": original_torch, "torch_cuda": original_cuda}, flush=True)
228
+
229
+
230
+ def native_build_options():
231
+ """Derive the compiler and supported SM target without replacing CUDA.
232
+
233
+ 推导编译器与支持的 SM 目标,不替换 CUDA;旧编译器为新显卡保留 PTX。
234
+ """
235
+ import torch
236
+ from importlib import metadata
237
+ nvcc = shutil.which("nvcc") or "/usr/local/cuda/bin/nvcc"
238
+ if not os.path.isfile(nvcc):
239
+ raise RuntimeError("CUDA compiler nvcc is missing / 缺少 CUDA 编译器 nvcc")
240
+ version_output = subprocess.check_output([nvcc, "--version"], text=True)
241
+ match = re.search(r"release (\d+)\.(\d+)", version_output)
242
+ if match is None:
243
+ raise RuntimeError("Cannot determine CUDA compiler version / 无法确定 CUDA 编译器版本")
244
+ compiler_version = tuple(int(value) for value in match.groups())
245
+ targets_output = subprocess.check_output([nvcc, "--list-gpu-code"], text=True)
246
+ supported = sorted({int(value) for value in re.findall(r"sm_(\d+)", targets_output)})
247
+ major, minor = torch.cuda.get_device_capability(0)
248
+ target_sm = major * 10 + minor
249
+ supported_sm = max(value for value in supported if value <= target_sm)
250
+ architecture = str(target_sm) if target_sm in supported else str(supported_sm) + "-virtual"
251
+ # Pybind11 resolves its CMake package using the active interpreter.
252
+ # pybind11 的 CMake 包使用当前解释器定位;CUDA 不依赖 torch 扩展的 ABI。
253
+ print("WAPR_NATIVE_BUILD", {"nvcc": nvcc, "compiler_cuda": compiler_version,
254
+ "torch_cuda": torch.version.cuda, "gpu_sm": target_sm,
255
+ "cmake_architecture": architecture,
256
+ "pybind11": metadata.version("pybind11")}, flush=True)
257
+ return ["-DCMAKE_CUDA_COMPILER=" + nvcc, "-DCMAKE_CUDA_ARCHITECTURES=" + architecture]
258
+
259
+
260
+ def fetch_example():
261
+ """Fetch the public usage example from the package's GitHub source URL.
262
+
263
+ 从包内指定的 GitHub 地址下载公开用法示例,不执行下载的代码。
264
+ """
265
+ from wapr.resources import resource_root
266
+ urls = ["https://raw.githubusercontent.com/WangYuLin-SEU/WAPR/main/examples/02_one_category_one_instance.py",
267
+ "https://api.github.com/repos/WangYuLin-SEU/WAPR/contents/examples/02_one_category_one_instance.py?ref=main"]
268
+ destination = Path(resource_root()) / "examples" / "02_one_category_one_instance.py"
269
+ if not destination.is_file():
270
+ destination.parent.mkdir(parents=True, exist_ok=True)
271
+ content = None
272
+ for url in urls:
273
+ # Direct GitHub often works when an acceleration proxy returns 503.
274
+ # 加速代理返回 503 时,GitHub 直连常仍可用;只在包内执行回退。
275
+ for direct in (True, False):
276
+ opener = build_opener(ProxyHandler({})) if direct else build_opener()
277
+ request = Request(url, headers={"Accept": "application/vnd.github.raw+json",
278
+ "User-Agent": "WAPR-resource-loader"})
279
+ try:
280
+ with opener.open(request, timeout=20) as response:
281
+ content = response.read()
282
+ if b"WAPREstimator" not in content:
283
+ content = None
284
+ continue
285
+ break
286
+ except (HTTPError, URLError, TimeoutError, OSError) as error:
287
+ print("WAPR_GITHUB_RETRY", {"direct": direct, "error": type(error).__name__}, flush=True)
288
+ if content is not None:
289
+ break
290
+ if content is None:
291
+ raise RuntimeError("Cannot download the GitHub usage example / 无法下载 GitHub 用法示例")
292
+ if b"WAPREstimator" not in content:
293
+ raise RuntimeError("Unexpected GitHub example content / GitHub 示例内容异常")
294
+ destination.write_bytes(content)
295
+ print("WAPR_GITHUB_EXAMPLE", str(destination), flush=True)
296
+ return destination
297
+
298
+
299
+ if __name__ == "__main__":
300
+ parser = argparse.ArgumentParser(description="Prepare WAPR in the existing Python / 在当前 Python 准备 WAPR")
301
+ parser.add_argument("--feature", nargs="+", default=["core"], choices=[
302
+ "core", "dinov2", "det2d", "sam2", "sam3d", "roma", "qwen", "robot", "unipose9d", "compatible"])
303
+ parser.add_argument("--check", action="store_true", help="Inspect without installing / 只检查,不安装")
304
+ parser.add_argument("--yes", action="store_true", help="Explicitly approve shown package replacements / 明确同意包更换")
305
+ arguments = parser.parse_args()
306
+ approved = True if arguments.yes else None
307
+ for requested_feature in arguments.feature:
308
+ preparation = prepare_feature(requested_feature, allow_replacement=approved, check_only=arguments.check)
309
+ print("WAPR_PREPARATION", json.dumps(preparation, ensure_ascii=False), flush=True)
310
+ if preparation.get("status") in ("blocked", "declined", "failed", "approval_required", "restart_required", "partial", "needs_source", "dependencies_ready") and not arguments.check:
311
+ raise SystemExit(1)
wapr/bop.py ADDED
@@ -0,0 +1,217 @@
1
+ # Author: Yulin Wang (yulinwang@seu.edu.cn)
2
+ # School of Mechanical Engineering, Southeast University, China
3
+ # Copyright (c) 2026 Yulin Wang. All rights reserved, except as granted under LICENSE.
4
+ # SPDX-License-Identifier: LGPL-2.1-only
5
+ # 作者与版权人:Yulin Wang;使用、修改与再分发须遵守项目 LICENSE。
6
+ # Third-party portions retain their original notices and terms; see THIRD_PARTY_NOTICES.txt.
7
+
8
+ """Read BOP frames, CADs and published detections; write metric pose results.
9
+
10
+ 读取 BOP 帧、CAD 与已公布检测,输出米制位姿对应的 BOP 结果。
11
+ Dataset/frame selection and inference loops remain in the caller's recipe.
12
+ 数据集、帧的选择与推理循环由调用方配置明确给出。
13
+ """
14
+
15
+ import json
16
+ import os
17
+
18
+ import cv2
19
+ import numpy as np
20
+ import trimesh
21
+
22
+ from wapr.estimator import bbox_model_center
23
+
24
+
25
+ def read_image(path):
26
+ """Read stored channels and bit depth unchanged. / 保留存储通道与位深读取图像。"""
27
+ image = cv2.imread(path, cv2.IMREAD_UNCHANGED)
28
+ if image is None:
29
+ raise FileNotFoundError(path)
30
+ return image
31
+
32
+
33
+ def test_split(dataset):
34
+ """Resolve the BOP depth-bearing test split. / 确定含深度图的 BOP 测试划分。"""
35
+ if str(dataset).lower() in ("tless", "hb"):
36
+ return "test_primesense"
37
+ return "test"
38
+
39
+
40
+ def first_image(scene_dir, folder, im_id):
41
+ """Find a six-digit BOP frame stem. / 查找六位编号的 BOP 帧文件。"""
42
+ stem = "%06d" % int(im_id)
43
+ for ext in (".png", ".jpg", ".tif", ".tiff"):
44
+ path = os.path.join(scene_dir, folder, stem + ext)
45
+ if os.path.isfile(path):
46
+ return path
47
+ return None
48
+
49
+
50
+ def load_bop_rgbd(bop_path, dataset, scene_id, im_id):
51
+ """Return RGB, float32 depth in meters, and 3×3 pixel intrinsics.
52
+
53
+ Stored depth times depth_scale is millimeters. Grayscale is repeated to RGB.
54
+ 返回 RGB、float32 米制深度与 3×3 像素内参。
55
+ 存储深度乘 depth_scale 得到毫米;灰度图重复为三个通道。
56
+ """
57
+ scene_dir = os.path.join(bop_path, dataset, test_split(dataset), "%06d" % int(scene_id))
58
+ with open(os.path.join(scene_dir, "scene_camera.json"), "r") as stream:
59
+ cam = json.load(stream)[str(int(im_id))]
60
+ K = np.asarray(cam["cam_K"], dtype=np.float32).reshape(3, 3)
61
+ depth_scale = float(cam.get("depth_scale", 1.0))
62
+ rgb_path = first_image(scene_dir, "rgb", im_id) or first_image(scene_dir, "gray", im_id)
63
+ depth_path = first_image(scene_dir, "depth", im_id)
64
+ rgb = read_image(rgb_path)
65
+ if rgb.ndim == 2:
66
+ rgb = np.repeat(rgb[..., None], 3, axis=2)
67
+ if rgb.ndim == 3 and rgb.shape[-1] == 3:
68
+ rgb = cv2.cvtColor(rgb, cv2.COLOR_BGR2RGB)
69
+ rgb = np.asarray(rgb[..., :3])
70
+ depth = np.asarray(read_image(depth_path), dtype=np.float32)
71
+ if depth.ndim == 3:
72
+ depth = depth[..., 0]
73
+ depth_m = depth * depth_scale / 1000.0
74
+ return rgb, depth_m, K
75
+
76
+
77
+ def model_ids(bop_path, dataset):
78
+ """List sorted ids from models/obj_######.ply. / 按文件名列出排序后的 CAD 编号。"""
79
+ models = os.path.join(bop_path, dataset, "models")
80
+ found = []
81
+ for name in sorted(os.listdir(models)):
82
+ if name.startswith("obj_") and name.endswith(".ply"):
83
+ found.append(int(name[len("obj_"):-len(".ply")]))
84
+ return found
85
+
86
+
87
+ def load_mesh_m(bop_path, dataset, obj_id):
88
+ """Convert BOP millimeter geometry/diameter to meters and center the mesh.
89
+
90
+ model_center preserves the original CAD origin used by output poses.
91
+ 将 BOP 毫米网格与直径换为米,并居中网格。
92
+ model_center 保留输出位姿对应的原始 CAD 原点偏移。
93
+ """
94
+ models = os.path.join(bop_path, dataset, "models")
95
+ ply = os.path.join(models, "obj_%06d.ply" % int(obj_id))
96
+ loaded = trimesh.load(ply, force="mesh", process=False)
97
+ if not isinstance(loaded, trimesh.Trimesh):
98
+ raise TypeError("expected one mesh")
99
+ mesh = loaded.copy()
100
+ mesh.apply_scale(0.001)
101
+ center = bbox_model_center(mesh.vertices)
102
+ mesh.vertices = np.asarray(mesh.vertices, dtype=np.float32) - center.reshape(1, 3)
103
+ mesh.metadata["model_center"] = center
104
+ mesh.metadata["name"] = "obj_%06d" % int(obj_id)
105
+ with open(os.path.join(models, "models_info.json"), "r") as stream:
106
+ info = json.load(stream)[str(int(obj_id))]
107
+ diameter_m = float(info["diameter"]) / 1000.0
108
+ return mesh, diameter_m
109
+
110
+
111
+ def rle_to_mask(rle):
112
+ """Decode supplied COCO RLE, returning None for unavailable segmentation.
113
+
114
+ 解码给定的 COCO RLE;分割不可用时返回 None,供调用方使用检测框。
115
+ """
116
+ if not rle:
117
+ return None
118
+ counts = rle.get("counts")
119
+ size = rle.get("size")
120
+ if counts is None or size is None:
121
+ return None
122
+ if isinstance(counts, str):
123
+ try:
124
+ from pycocotools import mask as mask_utils
125
+ return np.asarray(mask_utils.decode(rle), dtype=np.uint8)
126
+ except Exception:
127
+ return None
128
+ binary = np.zeros(int(np.prod(size)), dtype=np.uint8)
129
+ start = 0
130
+ for i in range(len(counts) - 1):
131
+ start += int(counts[i])
132
+ end = start + int(counts[i + 1])
133
+ binary[start:end] = (i + 1) % 2
134
+ return binary.reshape(int(size[0]), int(size[1]), order="F")
135
+
136
+
137
+ def load_detections(path, score_thr=0.0):
138
+ """Group published predicted boxes/masks by (scene_id, im_id), without GT.
139
+
140
+ 按 (scene_id, im_id) 组织公布的预测框与掩码,不读取真值。
141
+ """
142
+ with open(path, "r") as stream:
143
+ raw = json.load(stream)
144
+ if isinstance(raw, list):
145
+ items = raw
146
+ else:
147
+ items = []
148
+ for value in raw.values():
149
+ if isinstance(value, list):
150
+ items.extend(value)
151
+ elif isinstance(value, dict):
152
+ items.append(value)
153
+ out = {}
154
+ for det in items:
155
+ scene_id = int(det.get("scene_id", det.get("sceneId", -1)))
156
+ im_id = int(det.get("image_id", det.get("im_id", det.get("imageId", -1))))
157
+ obj_id = int(det.get("category_id", det.get("obj_id", det.get("categoryId", -1))))
158
+ score = float(det.get("score", det.get("confidence", 0.0)))
159
+ if score < float(score_thr):
160
+ continue
161
+ bbox = det.get("bbox", det.get("bbox_est"))
162
+ if bbox is None:
163
+ continue
164
+ mask = None
165
+ if det.get("segmentation") is not None:
166
+ mask = rle_to_mask(det["segmentation"])
167
+ out.setdefault((scene_id, im_id), []).append(
168
+ {"obj_id": obj_id, "bbox_xywh": [float(x) for x in bbox], "score_2d": score, "mask": mask}
169
+ )
170
+ return out
171
+
172
+
173
+ def keep_top_per_class(records, max_per_class):
174
+ """Keep requested per-class counts in descending 2D score order.
175
+
176
+ 按 2D 分数降序保留各类别请求的实例数量。
177
+ """
178
+ grouped = {}
179
+ for record in records:
180
+ grouped.setdefault(int(record["obj_id"]), []).append(record)
181
+ kept = []
182
+ for obj_id, rows in grouped.items():
183
+ ordered = sorted(rows, key=lambda row: float(row["score_2d"]), reverse=True)
184
+ cap = max_per_class.get(obj_id)
185
+ kept.extend(ordered if cap is None else ordered[:int(cap)])
186
+ return kept
187
+
188
+
189
+ def target_frames(bop_path, dataset):
190
+ """Read unique official target frame keys, without object counts.
191
+
192
+ 读取官方目标文件中去重的帧编号,不向检测流程提供标注物体数量。
193
+ """
194
+ path = os.path.join(bop_path, dataset, "test_targets_bop19.json")
195
+ with open(path, "r") as stream:
196
+ targets = json.load(stream)
197
+ return sorted({(int(row["scene_id"]), int(row["im_id"])) for row in targets})
198
+
199
+
200
+ def write_csv(path, rows):
201
+ """Write row-major R, millimeter t, score_6d and whole-frame seconds.
202
+
203
+ All rows of an image receive the caller's same whole-frame time.
204
+ 写出按行展开的 R、毫米 t、score_6d 与整帧秒数。
205
+ 同一图像的各行使用调用方给出的同一个整帧耗时。
206
+ """
207
+ os.makedirs(os.path.dirname(path) or ".", exist_ok=True)
208
+ with open(path, "w") as stream:
209
+ stream.write("scene_id,im_id,obj_id,score,R,t,time\n")
210
+ for row in rows:
211
+ R = np.asarray(row["R"], dtype=np.float64).reshape(-1)
212
+ t_mm = np.asarray(row["t_m"], dtype=np.float64).reshape(3) * 1000.0
213
+ stream.write("%d,%d,%d,%.8f,%s,%s,%.6f\n" % (
214
+ int(row["scene_id"]), int(row["im_id"]), int(row["obj_id"]),
215
+ float(row["score_6d"]), " ".join("%.8f" % float(x) for x in R),
216
+ " ".join("%.8f" % float(x) for x in t_mm), float(row["time"]),
217
+ ))