ultralytics_compat.py 3.2 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586
  1. """
  2. ultralytics 导入兼容层。
  3. 背景:本框架根目录下有一个同名的 `ultralytics/` 目录(YOLO 源码仓库根,
  4. 顶层无真正可 import 的包语义,或嵌套 `ultralytics/ultralytics/`)。
  5. 运行脚本通常把框架根插到 sys.path 首位(为 `import config` / `from modules import ...`),
  6. 这会让该目录以"命名空间包"形式遮蔽已(pip -e)安装、真正的 ultralytics 包,
  7. 导致 `from ultralytics import YOLO` 取到空包而报 ImportError。
  8. 另外:若进程 cwd 恰好是框架根,sys.path[0] 常为 ''(空串),旧逻辑
  9. `if not p: return False` 会漏判遮蔽,必须把 '' 解析成真实 cwd。
  10. 解决:导入前,把 sys.path 中"含同名 ultralytics/ 子目录但没有 __init__.py"的遮蔽条目
  11. 临时移除;若 sys.modules 里已有半残的 ultralytics 命名空间包也一并清掉再 import;
  12. 导入后再把移除的条目放回 sys.path 末尾。
  13. """
  14. from __future__ import annotations
  15. import os
  16. import sys
  17. def _framework_root():
  18. return os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
  19. def _resolve_path_entry(p: str) -> str:
  20. """把 sys.path 条目解析成绝对目录。'' / '.' → cwd。"""
  21. if p is None:
  22. return ""
  23. if p == "" or p == ".":
  24. return os.path.abspath(os.getcwd())
  25. return os.path.abspath(p)
  26. def _is_shadowing(p: str) -> bool:
  27. """该 path 项下是否有同名的、非真正包的 ultralytics/ 目录(命名空间遮蔽源)。"""
  28. root = _resolve_path_entry(p)
  29. if not root:
  30. return False
  31. d = os.path.join(root, "ultralytics")
  32. # 顶层 ultralytics/ 无 __init__.py → 命名空间包遮蔽
  33. if os.path.isdir(d) and not os.path.exists(os.path.join(d, "__init__.py")):
  34. return True
  35. return False
  36. def import_yolo():
  37. """返回 ultralytics.YOLO,规避框架根下同名目录的命名空间包遮蔽。"""
  38. # 1) 去掉遮蔽 path
  39. shadow = [p for p in list(sys.path) if _is_shadowing(p)]
  40. for p in shadow:
  41. try:
  42. sys.path.remove(p)
  43. except ValueError:
  44. pass
  45. # 2) 清掉可能已缓存的半残 ultralytics 模块(命名空间包 / 空包)
  46. doomed = [k for k in list(sys.modules) if k == "ultralytics" or k.startswith("ultralytics.")]
  47. # 仅当当前 ultralytics 看起来不可用(无 YOLO)时才清
  48. need_clear = False
  49. mod = sys.modules.get("ultralytics")
  50. if mod is None:
  51. need_clear = False
  52. else:
  53. if not hasattr(mod, "YOLO"):
  54. need_clear = True
  55. else:
  56. # 即便有 YOLO,若 __file__ 指向框架根下命名空间也不可靠——有 YOLO 就信
  57. need_clear = False
  58. if need_clear:
  59. for k in doomed:
  60. sys.modules.pop(k, None)
  61. try:
  62. from ultralytics import YOLO # noqa: WPS433
  63. if not hasattr(YOLO, "__call__") and not callable(YOLO):
  64. # 极端兜底:再清一次强刷
  65. for k in list(sys.modules):
  66. if k == "ultralytics" or k.startswith("ultralytics."):
  67. sys.modules.pop(k, None)
  68. from ultralytics import YOLO # noqa: WPS433
  69. return YOLO
  70. finally:
  71. for p in shadow:
  72. if p not in sys.path:
  73. sys.path.append(p)