Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions src/one_dragon/base/config/basic_model_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@
DEFAULT_OCR_MODEL_NAME,
PPOCRV6_MODEL_NAME,
get_final_file_list,
get_ocr_download_url_gitee,
get_ocr_download_url_cnb,
get_ocr_download_url_github,
get_ocr_model_dir,
)
Expand Down Expand Up @@ -48,8 +48,8 @@ def get_ocr_opts() -> list[ConfigItem]:
param = CommonDownloaderParam(
save_file_path=model_dir,
save_file_name=zip_file_name,
cnb_release_download_url=get_ocr_download_url_cnb(model),
github_release_download_url=get_ocr_download_url_github(model),
gitee_release_download_url=get_ocr_download_url_gitee(model),
check_existed_list=get_final_file_list(model),
)
config_list.append(
Expand Down
10 changes: 6 additions & 4 deletions src/one_dragon/base/matcher/ocr/onnx_ocr_matcher.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,7 @@
DEFAULT_OCR_MODEL_NAME: str = 'ppocrv5'
PPOCRV6_MODEL_NAME: str = 'ppocrv6'
GITHUB_DOWNLOAD_URL: str = 'https://github.com/OneDragon-Anything/OneDragon-Env/releases/download'
GITEE_DOWNLOAD_URL: str = 'https://gitee.com/OneDragon-Anything/OneDragon-Env/releases/download'
CNB_DOWNLOAD_URL: str = 'https://cnb.cool/OneDragon-Anything/OneDragon-Env/-/releases/download'


def get_ocr_model_dir(ocr_model_name: str) -> str:
Expand All @@ -31,8 +31,8 @@ def get_ocr_download_url_github(ocr_model_name: str) -> str:
return get_ocr_download_url(GITHUB_DOWNLOAD_URL, ocr_model_name)


def get_ocr_download_url_gitee(ocr_model_name: str) -> str:
return get_ocr_download_url(GITEE_DOWNLOAD_URL, ocr_model_name)
def get_ocr_download_url_cnb(ocr_model_name: str) -> str:
return get_ocr_download_url(CNB_DOWNLOAD_URL, ocr_model_name)


def get_ocr_download_url(website: str, ocr_model_name: str) -> str:
Expand Down Expand Up @@ -151,8 +151,8 @@ def __init__(self, ocr_param: OnnxOcrParam | None = None):
param = CommonDownloaderParam(
save_file_path=ocr_param.models_dir,
save_file_name=f'{ocr_param.ocr_model_name}.zip',
cnb_release_download_url=get_ocr_download_url_cnb(ocr_param.ocr_model_name),
github_release_download_url=get_ocr_download_url_github(ocr_param.ocr_model_name),
gitee_release_download_url=get_ocr_download_url_gitee(ocr_param.ocr_model_name),
mirror_chan_download_url='',
check_existed_list=get_final_file_list(ocr_param.ocr_model_name)
)
Expand Down Expand Up @@ -199,6 +199,7 @@ def _rect_from_anchor(anchor_position) -> tuple[int, int, int, int] | None:

def init_model(
self,
download_by_cnb: bool = False,
download_by_github: bool = True,
download_by_gitee: bool = False,
download_by_mirror_chan: bool = False,
Expand All @@ -216,6 +217,7 @@ def init_model(

# 先检查模型文件和下载模型
done: bool = self.download(
download_by_cnb=download_by_cnb,
download_by_github=download_by_github,
download_by_gitee=download_by_gitee,
download_by_mirror_chan=download_by_mirror_chan,
Expand Down
53 changes: 35 additions & 18 deletions src/one_dragon/base/web/common_downloader.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,18 +14,20 @@ def __init__(
github_release_download_url: str | None = None,
gitee_release_download_url: str | None = None,
mirror_chan_download_url: str | None = None,
cnb_release_download_url: str | None = None,
check_existed_list: list[str] | None = None,
unzip_dir_path: str | None = None,
):
"""
一个通用下载器 可提供3个下载源 并检查文件是否存在 如果存在则不进行下载
一个通用下载器 可提供多个下载源 并检查文件是否存在 如果存在则不进行下载

Args:
save_file_path (str): 文件保存的路径
save_file_name (str): 文件保存的名称
github_release_download_url (Optional[str], optional): Github Release下载地址. Defaults to None.
gitee_release_download_url (Optional[str], optional): Gitee Release下载地址. Defaults to None.
mirror_chan_download_url (Optional[str], optional): Mirror酱下载地址. Defaults to None.
cnb_release_download_url (Optional[str], optional): CNB Release下载地址. Defaults to None.
check_existed_list (Optional[list[str]], optional): 需要检查文件是否存在的列表 完整路径的列表. Defaults to None.
unzip_dir_path (Optional[str], optional): 解压目录路径,如果为None则解压到save_file_path. Defaults to None.
"""
Expand All @@ -34,6 +36,7 @@ def __init__(
self.github_release_download_url: str | None = github_release_download_url
self.gitee_release_download_url: str | None = gitee_release_download_url
self.mirror_chan_download_url: str | None = mirror_chan_download_url
self.cnb_release_download_url: str | None = cnb_release_download_url
self.check_existed_list: list[str] = [] if check_existed_list is None else check_existed_list
self.unzip_dir_path: str | None = unzip_dir_path

Expand All @@ -54,6 +57,7 @@ def __init__(

def download(
self,
download_by_cnb: bool = False,
download_by_github: bool = True,
download_by_gitee: bool = False,
download_by_mirror_chan: bool = False,
Expand All @@ -66,27 +70,40 @@ def download(
if skip_if_existed and self.is_file_existed():
return True

download_url: str = ''
# 按优先级顺序尝试启用的下载源:CNB → GitHub → Gitee → Mirror酱
# 某个源失败后自动尝试下一个 全部失败才返回 False
candidates: list[tuple[str, bool]] = []
if download_by_cnb and self.param.cnb_release_download_url is not None:
candidates.append((self.param.cnb_release_download_url, False))
if download_by_github and self.param.github_release_download_url is not None:
if ghproxy_url is not None:
download_url=f'{ghproxy_url}/{self.param.github_release_download_url}'
else:
download_url = self.param.github_release_download_url
elif download_by_gitee and self.param.gitee_release_download_url is not None:
download_url = self.param.gitee_release_download_url
elif download_by_mirror_chan and self.param.mirror_chan_download_url is not None:
download_url = self.param.mirror_chan_download_url

if download_url == '':
candidates.append((self.param.github_release_download_url, True))
if download_by_gitee and self.param.gitee_release_download_url is not None:
candidates.append((self.param.gitee_release_download_url, False))
if download_by_mirror_chan and self.param.mirror_chan_download_url is not None:
candidates.append((self.param.mirror_chan_download_url, False))

if not candidates:
log.error('没有指定下载方法或对应的下载地址')
return False

return http_utils.download_file(
download_url=download_url,
save_file_path=os.path.join(self.param.save_file_path, self.param.save_file_name),
proxy=proxy_url,
progress_signal=progress_signal,
progress_callback=progress_callback)
for download_url, use_ghproxy in candidates:
if use_ghproxy and ghproxy_url is not None:
download_url = f'{ghproxy_url}/{download_url}'
try:
if http_utils.download_file(
download_url=download_url,
save_file_path=os.path.join(self.param.save_file_path, self.param.save_file_name),
proxy=proxy_url,
progress_signal=progress_signal,
progress_callback=progress_callback,
):
return True
except Exception:
log.error(f'下载源失败: {download_url}', exc_info=True)
continue
log.warning(f'下载源失败 尝试下一个: {download_url}')

return False

def is_file_existed(self) -> bool:
"""
Expand Down
2 changes: 2 additions & 0 deletions src/one_dragon/base/web/zip_downloader.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,7 @@ def __init__(

def download(
self,
download_by_cnb: bool = False,
download_by_github: bool = True,
download_by_gitee: bool = False,
download_by_mirror_chan: bool = False,
Expand All @@ -40,6 +41,7 @@ def download(
for i in range(2):
download_result = CommonDownloader.download(
self,
download_by_cnb=download_by_cnb,
download_by_github=download_by_github,
download_by_gitee=download_by_gitee,
download_by_mirror_chan=download_by_mirror_chan,
Expand Down
41 changes: 29 additions & 12 deletions src/one_dragon/yolo/onnx_model_loader.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,10 +22,12 @@ def __init__(self,
personal_proxy: str | None = '',
gpu: bool = False,
backup_model_name: str | None = None,
backup_model_download_url: str | None = None,
):
self.model_name: str = model_name
self.backup_model_name: str = backup_model_name # 备用模型 默认在本地一定有的模型 在新模型无法下载使用时使用
self.model_download_url: str = model_download_url # 模型下载地址
self.backup_model_download_url: str | None = backup_model_download_url # 备用模型下载地址 主地址失败时使用
self.model_parent_dir_path: str = model_parent_dir_path
self.model_dir_path = os.path.join(self.model_parent_dir_path, self.model_name)
self.gh_proxy: bool = gh_proxy
Expand Down Expand Up @@ -71,18 +73,29 @@ def check_model_exists(self) -> bool:
def download_model(self) -> bool:
"""
下载模型
主下载地址失败后 自动尝试备用下载地址
:return: 是否成功下载模型
"""
if not os.path.exists(self.model_dir_path):
os.mkdir(self.model_dir_path)

download_url = f'{self.model_download_url}/{self.model_name}.zip'
if self.personal_proxy is not None and len(self.personal_proxy) > 0:
os.environ['http_proxy'] = self.personal_proxy
os.environ['https_proxy'] = self.personal_proxy
elif self.gh_proxy:
download_url = f'{self.gh_proxy_url}/{self.model_download_url}/{self.model_name}.zip'
log.info('开始下载 %s %s', self.model_name, download_url)

# 主地址优先(CNB 直连),备用地址(GitHub)可走 ghproxy 加速;
# 通过域名判断,GitHub 地址自动加 ghproxy 前缀,CNB 等地址直连
download_urls: list[str] = []
primary_url = f'{self.model_download_url}/{self.model_name}.zip'
if self.gh_proxy and 'github.com' in self.model_download_url:
primary_url = f'{self.gh_proxy_url}/{primary_url}'
download_urls.append(primary_url)
if self.backup_model_download_url:
backup_url = f'{self.backup_model_download_url}/{self.model_name}.zip'
if self.gh_proxy and 'github.com' in self.backup_model_download_url:
backup_url = f'{self.gh_proxy_url}/{backup_url}'
download_urls.append(backup_url)

zip_file_path = os.path.join(self.model_dir_path, f'{self.model_name}.zip')
last_log_time = time.time()

Expand All @@ -96,14 +109,18 @@ def log_download_progress(block_num, block_size, total_size):
progress = downloaded / total_size_mb * 100
log.info(f"正在下载 {self.model_name}: {downloaded:.2f}/{total_size_mb:.2f} MB ({progress:.2f}%)")

try:
_, _ = urllib.request.urlretrieve(download_url, zip_file_path, log_download_progress)
log.info('下载完成 %s', self.model_name)
self.unzip_model(zip_file_path)
return True
except Exception:
log.error('下载失败模型失败', exc_info=True)
return False
for download_url in download_urls:
log.info('开始下载 %s %s', self.model_name, download_url)
try:
_, _ = urllib.request.urlretrieve(download_url, zip_file_path, log_download_progress)
log.info('下载完成 %s', self.model_name)
self.unzip_model(zip_file_path)
return True
except Exception:
log.error(f'下载模型失败: {download_url}', exc_info=True)
continue

return False

def unzip_model(self, zip_file_path: str):
"""
Expand Down
5 changes: 5 additions & 0 deletions src/one_dragon/yolo/yolo_utils.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
_GITHUB_YOLO_RELEASE_BASE = 'https://github.com/OneDragon-Anything/OneDragon-YOLO/releases/download'
_GITEE_YOLO_RELEASE_BASE = 'https://gitee.com/OneDragon-Anything/OneDragon-YOLO/releases/download'
_CNB_YOLO_RELEASE_BASE = 'https://cnb.cool/OneDragon-Anything/OneDragon-YOLO/-/releases/download'


def get_github_model_download_url(release_tag: str) -> str:
Expand All @@ -8,3 +9,7 @@ def get_github_model_download_url(release_tag: str) -> str:

def get_gitee_model_download_url(release_tag: str) -> str:
return f'{_GITEE_YOLO_RELEASE_BASE}/{release_tag}'


def get_cnb_model_download_url(release_tag: str) -> str:
return f'{_CNB_YOLO_RELEASE_BASE}/{release_tag}'
6 changes: 5 additions & 1 deletion src/one_dragon/yolo/yolov8_onnx_cls.py
Original file line number Diff line number Diff line change
Expand Up @@ -59,11 +59,14 @@ def __init__(self,
personal_proxy: Optional[str] = None,
gpu: bool = False,
backup_model_name: Optional[str] = None,
backup_model_download_url: str | None = None,
keep_result_seconds: float = 2,
):
"""
:param model_name: 模型名称 在根目录下会有一个以模型名称创建的子文件夹
:param model_parent_dir_path: 放置所有模型的根目录
:param model_download_url: 模型下载地址 失败后使用备用地址
:param backup_model_download_url: 备用模型下载地址
:param gpu: 是否启用GPU加速
:param keep_result_seconds: 保留多长时间的识别结果
"""
Expand All @@ -76,7 +79,8 @@ def __init__(self,
gh_proxy_url=gh_proxy_url,
personal_proxy=personal_proxy,
gpu=gpu,
backup_model_name=backup_model_name
backup_model_name=backup_model_name,
backup_model_download_url=backup_model_download_url
)

self.keep_result_seconds: float = keep_result_seconds # 保留识别结果的秒数
Expand Down
6 changes: 5 additions & 1 deletion src/one_dragon/yolo/yolov8_onnx_det.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,13 +23,16 @@ def __init__(self,
personal_proxy: Optional[str] = None,
gpu: bool = False,
backup_model_name: Optional[str] = None,
backup_model_download_url: str | None = None,
keep_result_seconds: float = 2
):
"""
yolov8 detect 导出 onnx 后使用
参考自 https://github.com/ibaiGorordo/ONNX-YOLOv8-Object-Detection
:param model_name: 模型名称 在根目录下会有一个以模型名称创建的子文件夹
:param backup_model_name: 备用模型名称 通常是上一个版本的模型 在新版本模型无法下载时兜底使用
:param model_download_url: 模型下载地址 失败后使用备用地址
:param backup_model_download_url: 备用模型下载地址
:param model_parent_dir_path: 放置所有模型的根目录
:param gpu: 是否启用GPU运算
:param keep_result_seconds: 保留多长时间的识别结果
Expand All @@ -43,7 +46,8 @@ def __init__(self,
gh_proxy_url=gh_proxy_url,
personal_proxy=personal_proxy,
gpu=gpu,
backup_model_name=backup_model_name
backup_model_name=backup_model_name,
backup_model_download_url=backup_model_download_url
)

self.keep_result_seconds: float = keep_result_seconds # 保留识别结果的秒数
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,8 @@ def run(self):
"""
try:
result = self.downloader.download(
download_by_cnb=True,
download_by_github=True,
ghproxy_url=self.ctx.env_config.gh_proxy_url if self.ctx.env_config.is_gh_proxy else None,
proxy_url=self.ctx.env_config.personal_proxy if self.ctx.env_config.is_personal_proxy else None,
skip_if_existed=False,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,10 @@

from one_dragon.utils import yolo_config_utils
from one_dragon.yolo.detect_utils import DetectFrameResult, DetectObjectResult
from one_dragon.yolo.yolo_utils import get_github_model_download_url
from one_dragon.yolo.yolo_utils import (
get_cnb_model_download_url,
get_github_model_download_url,
)
from one_dragon.yolo.yolov8_onnx_det import Yolov8Detector
from zzz_od.config.model_config import YOLO_RELEASE_TAG

Expand Down Expand Up @@ -41,7 +44,8 @@ def __init__(self,
model_name=model_name,
backup_model_name=backup_model_name,
model_parent_dir_path=yolo_config_utils.get_model_category_dir('lost_void_det'),
model_download_url=get_github_model_download_url(YOLO_RELEASE_TAG),
model_download_url=get_cnb_model_download_url(YOLO_RELEASE_TAG),
backup_model_download_url=get_github_model_download_url(YOLO_RELEASE_TAG),
gh_proxy=gh_proxy,
gh_proxy_url=gh_proxy_url,
personal_proxy=personal_proxy,
Expand Down
10 changes: 5 additions & 5 deletions src/zzz_od/config/model_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,13 +5,13 @@
from one_dragon.base.web.common_downloader import CommonDownloaderParam
from one_dragon.utils import yolo_config_utils
from one_dragon.yolo.yolo_utils import (
get_gitee_model_download_url,
get_cnb_model_download_url,
get_github_model_download_url,
)

YOLO_RELEASE_TAG = 'zzz_model'
YOLO_CNB_MODEL_DOWNLOAD_URL = get_cnb_model_download_url(YOLO_RELEASE_TAG)
YOLO_GITHUB_MODEL_DOWNLOAD_URL = get_github_model_download_url(YOLO_RELEASE_TAG)
YOLO_GITEE_MODEL_DOWNLOAD_URL = get_gitee_model_download_url(YOLO_RELEASE_TAG)

_DEFAULT_FLASH_CLASSIFIER = 'yolov8n-640-flash-20250921'
_BACKUP_FLASH_CLASSIFIER = 'yolov8n-640-flash-20250906'
Expand Down Expand Up @@ -135,8 +135,8 @@ def get_flash_classifier_opts() -> list[ConfigItem]:
param = CommonDownloaderParam(
save_file_path=model_dir,
save_file_name=zip_file_name,
cnb_release_download_url=f'{YOLO_CNB_MODEL_DOWNLOAD_URL}/{zip_file_name}',
github_release_download_url=f'{YOLO_GITHUB_MODEL_DOWNLOAD_URL}/{zip_file_name}',
gitee_release_download_url=f'{YOLO_GITEE_MODEL_DOWNLOAD_URL}/{zip_file_name}',
check_existed_list=[
os.path.join(model_dir, 'model.onnx'),
os.path.join(model_dir, 'labels.csv'),
Expand Down Expand Up @@ -167,8 +167,8 @@ def get_hollow_zero_event_opts() -> list[ConfigItem]:
param = CommonDownloaderParam(
save_file_path=model_dir,
save_file_name=zip_file_name,
cnb_release_download_url=f'{YOLO_CNB_MODEL_DOWNLOAD_URL}/{zip_file_name}',
github_release_download_url=f'{YOLO_GITHUB_MODEL_DOWNLOAD_URL}/{zip_file_name}',
gitee_release_download_url=f'{YOLO_GITEE_MODEL_DOWNLOAD_URL}/{zip_file_name}',
check_existed_list=[
os.path.join(model_dir, 'model.onnx'),
os.path.join(model_dir, 'labels.csv'),
Expand Down Expand Up @@ -199,8 +199,8 @@ def get_lost_void_det_opts() -> list[ConfigItem]:
param = CommonDownloaderParam(
save_file_path=model_dir,
save_file_name=zip_file_name,
cnb_release_download_url=f'{YOLO_CNB_MODEL_DOWNLOAD_URL}/{zip_file_name}',
github_release_download_url=f'{YOLO_GITHUB_MODEL_DOWNLOAD_URL}/{zip_file_name}',
gitee_release_download_url=f'{YOLO_GITEE_MODEL_DOWNLOAD_URL}/{zip_file_name}',
check_existed_list=[
os.path.join(model_dir, 'model.onnx'),
os.path.join(model_dir, 'labels.csv'),
Expand Down
Loading
Loading