feat(data-management): 实现无GUI数据集下载与自动分割脚本
- 通过MySQL按日期查询数据并分页获取图像信息 - 实现SimpleDownloader支持多线程批量下载,带重试和跳过已存在文件功能 - 依据类别自动创建目录并规范化文件命名,防止文件名冲突 - 支持自定义下载超时时间、最大重试次数、并发数等参数 - 下载结束后调用已有split_dataset模块按70/15/15比例分割为训练/验证/测试集 - 支持设置样本最小阈值、类别样本上限和标记样本不足类别 - 生成下载失败日志和运行摘要,方便任务监控及后续排查 - 命令行参数灵活配置数据源、日期范围、输出路径等运行参数
This commit is contained in:
@@ -0,0 +1,334 @@
|
|||||||
|
"""
|
||||||
|
无GUI数据集构建脚本
|
||||||
|
流程:按日期查询MySQL -> 下载到train按类别目录 -> 按70/15/15分割为train/val/test
|
||||||
|
"""
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
import random
|
||||||
|
import re
|
||||||
|
import sys
|
||||||
|
import threading
|
||||||
|
import time
|
||||||
|
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from datetime import datetime
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Dict, List, Optional, Tuple
|
||||||
|
|
||||||
|
import requests
|
||||||
|
|
||||||
|
# 允许从 data_management 直接运行
|
||||||
|
CURRENT_DIR = Path(__file__).resolve().parent
|
||||||
|
if str(CURRENT_DIR) not in sys.path:
|
||||||
|
sys.path.append(str(CURRENT_DIR))
|
||||||
|
|
||||||
|
from database import MySQLManager # noqa: E402
|
||||||
|
from split_dataset import split_dataset # noqa: E402
|
||||||
|
|
||||||
|
|
||||||
|
IMAGE_EXT_DEFAULT = ".jpg"
|
||||||
|
VALID_EXTS = {".jpg", ".jpeg", ".png", ".webp", ".bmp"}
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class DownloadResult:
|
||||||
|
success: bool
|
||||||
|
goods_name: str
|
||||||
|
image_url: str
|
||||||
|
saved_path: Optional[str] = None
|
||||||
|
error_message: Optional[str] = None
|
||||||
|
|
||||||
|
|
||||||
|
class SimpleDownloader:
|
||||||
|
"""无状态批量下载器(带重试和跳过)"""
|
||||||
|
|
||||||
|
def __init__(self, timeout: int = 30, max_retries: int = 2, workers: int = 8):
|
||||||
|
self.timeout = timeout
|
||||||
|
self.max_retries = max_retries
|
||||||
|
self.workers = workers
|
||||||
|
self._lock = threading.Lock()
|
||||||
|
self._seen_targets = set()
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _sanitize_name(name: str) -> str:
|
||||||
|
name = (name or "未知类别").strip()
|
||||||
|
name = re.sub(r"[\\/:*?\"<>|]", "_", name)
|
||||||
|
name = re.sub(r"\s+", "_", name)
|
||||||
|
return name[:80] if len(name) > 80 else name
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _infer_ext_from_url(image_url: str) -> str:
|
||||||
|
base = image_url.split("?")[0]
|
||||||
|
ext = os.path.splitext(base)[1].lower()
|
||||||
|
return ext if ext in VALID_EXTS else IMAGE_EXT_DEFAULT
|
||||||
|
|
||||||
|
def _build_file_path(self, train_root: Path, row: Dict) -> Path:
|
||||||
|
cls_name = self._sanitize_name(row.get("goods_name", "未知类别"))
|
||||||
|
class_dir = train_root / cls_name
|
||||||
|
class_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
|
create_time = row.get("create_time", "")
|
||||||
|
dt_tag = re.sub(r"[^0-9]", "", create_time) if create_time else ""
|
||||||
|
dt_tag = dt_tag[:14] if dt_tag else datetime.now().strftime("%Y%m%d%H%M%S")
|
||||||
|
|
||||||
|
data_id = row.get("id", "na")
|
||||||
|
ext = self._infer_ext_from_url(row.get("image_url", ""))
|
||||||
|
filename = f"img_{dt_tag}_{data_id}{ext}"
|
||||||
|
target = class_dir / filename
|
||||||
|
|
||||||
|
# 文件名冲突兜底
|
||||||
|
if target.exists():
|
||||||
|
suffix = 1
|
||||||
|
while True:
|
||||||
|
candidate = class_dir / f"img_{dt_tag}_{data_id}_{suffix}{ext}"
|
||||||
|
if not candidate.exists():
|
||||||
|
target = candidate
|
||||||
|
break
|
||||||
|
suffix += 1
|
||||||
|
|
||||||
|
return target
|
||||||
|
|
||||||
|
def _download_one(self, row: Dict, train_root: Path, skip_existing: bool) -> DownloadResult:
|
||||||
|
goods_name = row.get("goods_name", "未知类别")
|
||||||
|
image_url = row.get("image_url", "")
|
||||||
|
if not image_url:
|
||||||
|
return DownloadResult(False, goods_name, image_url, error_message="空图片URL")
|
||||||
|
|
||||||
|
target = self._build_file_path(train_root, row)
|
||||||
|
target_key = str(target)
|
||||||
|
|
||||||
|
with self._lock:
|
||||||
|
if target_key in self._seen_targets:
|
||||||
|
return DownloadResult(True, goods_name, image_url, saved_path=target_key)
|
||||||
|
self._seen_targets.add(target_key)
|
||||||
|
|
||||||
|
if skip_existing and target.exists():
|
||||||
|
return DownloadResult(True, goods_name, image_url, saved_path=str(target))
|
||||||
|
|
||||||
|
last_error = None
|
||||||
|
for _ in range(self.max_retries + 1):
|
||||||
|
try:
|
||||||
|
response = requests.get(image_url, timeout=self.timeout, stream=True)
|
||||||
|
response.raise_for_status()
|
||||||
|
with open(target, "wb") as f:
|
||||||
|
for chunk in response.iter_content(chunk_size=8192):
|
||||||
|
if chunk:
|
||||||
|
f.write(chunk)
|
||||||
|
return DownloadResult(True, goods_name, image_url, saved_path=str(target))
|
||||||
|
except Exception as e:
|
||||||
|
last_error = str(e)
|
||||||
|
time.sleep(0.3)
|
||||||
|
|
||||||
|
return DownloadResult(False, goods_name, image_url, error_message=last_error)
|
||||||
|
|
||||||
|
def download_batch(self, rows: List[Dict], train_root: Path, skip_existing: bool = True) -> Tuple[int, int, List[Dict]]:
|
||||||
|
success = 0
|
||||||
|
failed = 0
|
||||||
|
failures: List[Dict] = []
|
||||||
|
|
||||||
|
with ThreadPoolExecutor(max_workers=self.workers) as executor:
|
||||||
|
futures = [executor.submit(self._download_one, row, train_root, skip_existing) for row in rows]
|
||||||
|
for future in as_completed(futures):
|
||||||
|
result = future.result()
|
||||||
|
if result.success:
|
||||||
|
success += 1
|
||||||
|
else:
|
||||||
|
failed += 1
|
||||||
|
failures.append(
|
||||||
|
{
|
||||||
|
"goods_name": result.goods_name,
|
||||||
|
"image_url": result.image_url,
|
||||||
|
"error": result.error_message,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
return success, failed, failures
|
||||||
|
|
||||||
|
|
||||||
|
def parse_args() -> argparse.Namespace:
|
||||||
|
parser = argparse.ArgumentParser(description="按日期下载并自动分割70/15/15数据集")
|
||||||
|
|
||||||
|
parser.add_argument("--start-date", help="开始日期,格式 YYYY-MM-DD", default="2026-04-08")
|
||||||
|
parser.add_argument("--end-date", help="结束日期,格式 YYYY-MM-DD", default="2026-04-22")
|
||||||
|
parser.add_argument(
|
||||||
|
"--output-dir",
|
||||||
|
help="输出目录(每次手动指定)",
|
||||||
|
default=r"D:\MyProjects\PythonProjects\FoodClassifier\dataset\item\20260422",
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument("--datasource", default="食材", choices=["菜品", "食材", "半加工"], help="数据源")
|
||||||
|
parser.add_argument("--page-size", type=int, default=500, help="分页大小,默认500")
|
||||||
|
parser.add_argument("--workers", type=int, default=8, help="下载并发数,默认8")
|
||||||
|
parser.add_argument("--timeout", type=int, default=30, help="下载超时时间(秒)")
|
||||||
|
parser.add_argument("--max-retries", type=int, default=2, help="下载失败重试次数")
|
||||||
|
|
||||||
|
parser.add_argument("--train-ratio", type=float, default=0.70, help="训练集比例")
|
||||||
|
parser.add_argument("--val-ratio", type=float, default=0.15, help="验证集比例")
|
||||||
|
parser.add_argument("--test-ratio", type=float, default=0.15, help="测试集比例")
|
||||||
|
parser.add_argument("--seed", type=int, default=42, help="随机种子")
|
||||||
|
|
||||||
|
parser.add_argument("--min-samples", type=int, default=10, help="最小样本阈值")
|
||||||
|
parser.add_argument("--mark-insufficient", action="store_true", help="是否标记样本不足类别")
|
||||||
|
parser.add_argument("--max-total", type=int, default=120, help="单类上限,默认120")
|
||||||
|
|
||||||
|
parser.add_argument("--skip-existing", action="store_true", help="已存在文件直接跳过")
|
||||||
|
|
||||||
|
return parser.parse_args()
|
||||||
|
|
||||||
|
|
||||||
|
def parse_date(text: str) -> datetime:
|
||||||
|
return datetime.strptime(text, "%Y-%m-%d")
|
||||||
|
|
||||||
|
|
||||||
|
def fetch_all_rows(
|
||||||
|
mysql: MySQLManager,
|
||||||
|
datasource: str,
|
||||||
|
start_date: datetime,
|
||||||
|
end_date: datetime,
|
||||||
|
page_size: int,
|
||||||
|
) -> List[Dict]:
|
||||||
|
rows: List[Dict] = []
|
||||||
|
page = 1
|
||||||
|
|
||||||
|
while True:
|
||||||
|
page_data, total = mysql.query_images(
|
||||||
|
datasource=datasource,
|
||||||
|
start_date=start_date,
|
||||||
|
end_date=end_date,
|
||||||
|
goods_name=None,
|
||||||
|
equipment_code=None,
|
||||||
|
page=page,
|
||||||
|
page_size=page_size,
|
||||||
|
)
|
||||||
|
|
||||||
|
if not page_data:
|
||||||
|
break
|
||||||
|
|
||||||
|
rows.extend(page_data)
|
||||||
|
print(f"[Query] page={page}, fetched={len(page_data)}, accumulated={len(rows)}, total={total}")
|
||||||
|
|
||||||
|
if len(rows) >= total:
|
||||||
|
break
|
||||||
|
|
||||||
|
page += 1
|
||||||
|
|
||||||
|
return rows
|
||||||
|
|
||||||
|
|
||||||
|
def save_json(path: Path, obj: Dict):
|
||||||
|
path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
with open(path, "w", encoding="utf-8") as f:
|
||||||
|
json.dump(obj, f, ensure_ascii=False, indent=2)
|
||||||
|
|
||||||
|
|
||||||
|
def main():
|
||||||
|
args = parse_args()
|
||||||
|
|
||||||
|
start_date = parse_date(args.start_date)
|
||||||
|
end_date = parse_date(args.end_date)
|
||||||
|
if end_date < start_date:
|
||||||
|
raise ValueError("end_date 不能早于 start_date")
|
||||||
|
|
||||||
|
ratio_sum = args.train_ratio + args.val_ratio + args.test_ratio
|
||||||
|
if abs(ratio_sum - 1.0) > 1e-6:
|
||||||
|
raise ValueError(f"比例和必须为1.0,当前为 {ratio_sum}")
|
||||||
|
|
||||||
|
random.seed(args.seed)
|
||||||
|
|
||||||
|
output_dir = Path(args.output_dir).resolve()
|
||||||
|
train_dir = output_dir / "train"
|
||||||
|
val_dir = output_dir / "val"
|
||||||
|
test_dir = output_dir / "test"
|
||||||
|
logs_dir = output_dir / "logs"
|
||||||
|
|
||||||
|
train_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
val_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
test_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
logs_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
|
print("=" * 88)
|
||||||
|
print("数据集自动构建开始")
|
||||||
|
print(f"datasource={args.datasource}")
|
||||||
|
print(f"date_range=[{args.start_date}, {args.end_date}]")
|
||||||
|
print(f"output_dir={output_dir}")
|
||||||
|
print(f"ratio=train:{args.train_ratio}, val:{args.val_ratio}, test:{args.test_ratio}")
|
||||||
|
print("=" * 88)
|
||||||
|
|
||||||
|
mysql = MySQLManager()
|
||||||
|
rows = fetch_all_rows(mysql, args.datasource, start_date, end_date, args.page_size)
|
||||||
|
print(f"[Query] total rows fetched: {len(rows)}")
|
||||||
|
|
||||||
|
if not rows:
|
||||||
|
summary = {
|
||||||
|
"status": "empty",
|
||||||
|
"datasource": args.datasource,
|
||||||
|
"start_date": args.start_date,
|
||||||
|
"end_date": args.end_date,
|
||||||
|
"message": "查询结果为空,没有可下载数据",
|
||||||
|
}
|
||||||
|
save_json(logs_dir / "run_summary.json", summary)
|
||||||
|
print("没有数据,任务结束")
|
||||||
|
return
|
||||||
|
|
||||||
|
downloader = SimpleDownloader(
|
||||||
|
timeout=args.timeout,
|
||||||
|
max_retries=args.max_retries,
|
||||||
|
workers=args.workers,
|
||||||
|
)
|
||||||
|
|
||||||
|
success_count, failed_count, failures = downloader.download_batch(
|
||||||
|
rows=rows,
|
||||||
|
train_root=train_dir,
|
||||||
|
skip_existing=args.skip_existing,
|
||||||
|
)
|
||||||
|
|
||||||
|
print(f"[Download] success={success_count}, failed={failed_count}")
|
||||||
|
|
||||||
|
if failures:
|
||||||
|
save_json(logs_dir / "download_failures.json", {"failures": failures})
|
||||||
|
|
||||||
|
# 调用现有分割逻辑
|
||||||
|
split_dataset(
|
||||||
|
dataset_root=str(output_dir),
|
||||||
|
train_ratio=args.train_ratio,
|
||||||
|
val_ratio=args.val_ratio,
|
||||||
|
test_ratio=args.test_ratio,
|
||||||
|
min_samples_threshold=args.min_samples,
|
||||||
|
mark_insufficient=args.mark_insufficient,
|
||||||
|
max_total=args.max_total,
|
||||||
|
seed=args.seed,
|
||||||
|
)
|
||||||
|
|
||||||
|
summary = {
|
||||||
|
"status": "done",
|
||||||
|
"datasource": args.datasource,
|
||||||
|
"start_date": args.start_date,
|
||||||
|
"end_date": args.end_date,
|
||||||
|
"output_dir": str(output_dir),
|
||||||
|
"total_rows": len(rows),
|
||||||
|
"download_success": success_count,
|
||||||
|
"download_failed": failed_count,
|
||||||
|
"ratios": {
|
||||||
|
"train": args.train_ratio,
|
||||||
|
"val": args.val_ratio,
|
||||||
|
"test": args.test_ratio,
|
||||||
|
},
|
||||||
|
"seed": args.seed,
|
||||||
|
"page_size": args.page_size,
|
||||||
|
"workers": args.workers,
|
||||||
|
"skip_existing": args.skip_existing,
|
||||||
|
}
|
||||||
|
save_json(logs_dir / "run_summary.json", summary)
|
||||||
|
|
||||||
|
print("=" * 88)
|
||||||
|
print("数据集自动构建完成")
|
||||||
|
print(f"summary: {logs_dir / 'run_summary.json'}")
|
||||||
|
if failures:
|
||||||
|
print(f"failures: {logs_dir / 'download_failures.json'}")
|
||||||
|
print("=" * 88)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
Reference in New Issue
Block a user