Files
FoodClassifier/data_management/batch_download_all.py
T

299 lines
9.9 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
批量下载所有图片到dataset目录
自动按类别分类到对应文件夹
"""
import os
import sys
from datetime import datetime
from tqdm import tqdm
import requests
# 添加当前目录到路径
sys.path.append(os.path.dirname(os.path.abspath(__file__)))
from database import MySQLManager, SQLiteManager
from utils import ImageDownloader
from config import IMAGE_URL_PREFIX, DOWNLOAD_TIMEOUT, DOWNLOAD_THREADS
def sanitize_folder_name(name: str) -> str:
"""清理文件夹名称,移除非法字符"""
# 保留中文字符、字母、数字、括号、方括号、WP代码
import re
# 移除或替换非法字符
name = name.replace('/', '_').replace('\\', '_')
name = name.replace(':', '_').replace('*', '_')
name = name.replace('?', '_').replace('"', '_')
name = name.replace('<', '_').replace('>', '_')
name = name.replace('|', '_')
return name.strip()
def get_goods_code(goods_name: str) -> str:
"""从goods_name中提取WP代码(如果存在)"""
import re
# 匹配WP开头的代码,如 WP2025110700043
match = re.search(r'WP\d+', goods_name)
if match:
return match.group(0)
return ''
def batch_download_all(
datasource: str = '食材',
output_dir: str = None,
max_images: int = None,
skip_existing: bool = True
):
"""
批量下载所有图片
Args:
datasource: 数据源('食材'/'菜品'/'半加工'
output_dir: 输出目录,如果为None则使用dataset目录
max_images: 最大下载数量(用于测试),None表示下载全部
skip_existing: 是否跳过已存在的文件
"""
print("="*80)
print(f"批量下载工具 - 数据源: {datasource}")
print("="*80)
# 初始化数据库管理器
print("\n[1/5] 连接数据库...")
mysql_manager = MySQLManager()
sqlite_manager = SQLiteManager()
# 查询所有图片数据
print("\n[2/5] 查询所有图片数据...")
all_data = []
page = 1
page_size = 1000 # 每次查询1000条
while True:
data_list, total_count = mysql_manager.query_images(
datasource=datasource,
page=page,
page_size=page_size
)
if not data_list:
break
all_data.extend(data_list)
print(f" 已查询: {len(all_data)} / {total_count} 张图片")
if len(all_data) >= total_count:
break
if max_images and len(all_data) >= max_images:
all_data = all_data[:max_images]
break
page += 1
total_images = len(all_data)
print(f"\n✓ 共找到 {total_images} 张图片")
if total_images == 0:
print("没有找到图片,退出")
return
# 确定输出目录
if output_dir is None:
# 根据数据源确定目录
base_dir = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
if datasource == '食材':
output_dir = os.path.join(base_dir, 'dataset', 'WholeIngredientRecognition', 'train')
elif datasource == '菜品':
output_dir = os.path.join(base_dir, 'dataset', 'DishClassification', 'train')
elif datasource == '半加工':
output_dir = os.path.join(base_dir, 'dataset', 'ProcessedIngredientRecognition', 'train')
else:
output_dir = os.path.join(base_dir, 'dataset', datasource, 'train')
os.makedirs(output_dir, exist_ok=True)
print(f"\n[3/5] 输出目录: {output_dir}")
# 按类别分组
print("\n[4/5] 按类别分组...")
category_groups = {}
for data in all_data:
goods_name = data['goods_name']
# 清理文件夹名称
folder_name = sanitize_folder_name(goods_name)
if folder_name not in category_groups:
category_groups[folder_name] = []
category_groups[folder_name].append(data)
print(f"✓ 共 {len(category_groups)} 个类别")
# 显示类别统计
print("\n类别统计(前10个):")
sorted_categories = sorted(category_groups.items(), key=lambda x: len(x[1]), reverse=True)
for i, (category, items) in enumerate(sorted_categories[:10], 1):
print(f" {i}. {category}: {len(items)} 张")
# 下载图片
print(f"\n[5/5] 开始下载图片(并发数: {DOWNLOAD_THREADS}...")
success_count = 0
failed_count = 0
skipped_count = 0
error_messages = []
# 创建进度条
pbar = tqdm(total=total_images, desc="下载进度", unit="张")
def download_single(data: dict) -> tuple[bool, str]:
"""下载单张图片"""
try:
goods_name = data['goods_name']
folder_name = sanitize_folder_name(goods_name)
category_dir = os.path.join(output_dir, folder_name)
os.makedirs(category_dir, exist_ok=True)
# 构建文件名
filename = f"img_{data['id']}.jpg"
file_path = os.path.join(category_dir, filename)
# 检查文件是否已存在
if skip_existing and os.path.exists(file_path):
return True, "已存在"
# 下载图片
image_url = data['image_url']
if not image_url:
return False, "图片URL为空"
response = requests.get(image_url, timeout=DOWNLOAD_TIMEOUT, stream=True)
response.raise_for_status()
# 保存图片
with open(file_path, 'wb') as f:
for chunk in response.iter_content(chunk_size=8192):
if chunk:
f.write(chunk)
# 记录下载历史
sqlite_manager.add_download_record(
datasource=data['datasource'],
goods_id=data.get('goods_id'),
goods_name=goods_name,
image_url=image_url,
local_path=file_path,
file_size=os.path.getsize(file_path)
)
return True, None
except Exception as e:
return False, str(e)
# 使用线程池下载
from concurrent.futures import ThreadPoolExecutor, as_completed
with ThreadPoolExecutor(max_workers=DOWNLOAD_THREADS) as executor:
# 提交所有任务
futures = {executor.submit(download_single, data): data for data in all_data}
# 处理结果
for future in as_completed(futures):
data = futures[future]
try:
success, error = future.result()
if success:
if error == "已存在":
skipped_count += 1
else:
success_count += 1
else:
failed_count += 1
error_messages.append(f"{data['goods_name']} (ID: {data['id']}): {error}")
pbar.update(1)
except Exception as e:
failed_count += 1
error_messages.append(f"{data['goods_name']} (ID: {data['id']}): {str(e)}")
pbar.update(1)
pbar.close()
# 打印结果
print("\n" + "="*80)
print("下载完成!")
print("="*80)
print(f"成功: {success_count} 张")
print(f"跳过: {skipped_count} 张(已存在)")
print(f"失败: {failed_count} 张")
print(f"总计: {total_images} 张")
print(f"\n输出目录: {output_dir}")
if error_messages:
print(f"\n错误信息(前10条):")
for msg in error_messages[:10]:
print(f" - {msg}")
if len(error_messages) > 10:
print(f" ... 还有 {len(error_messages) - 10} 条错误")
# 关闭数据库连接
mysql_manager.close()
print("\n✓ 完成!")
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser(description='批量下载所有图片到dataset目录')
parser.add_argument('--datasource', type=str, default='食材',
choices=['食材', '菜品', '半加工', 'all'],
help='数据源(默认: 食材,使用 all 下载所有三种)')
parser.add_argument('--output', type=str, default=None,
help='输出目录(默认: dataset/{datasource}/train')
parser.add_argument('--max', type=int, default=None,
help='最大下载数量(用于测试)')
parser.add_argument('--no-skip', action='store_true',
help='不跳过已存在的文件(重新下载)')
args = parser.parse_args()
# 如果选择 all,则依次下载所有三种数据源
if args.datasource == 'all':
datasources = ['食材', '菜品', '半加工']
print("\n" + "="*80)
print("将依次下载所有三种数据源:")
print(" 1. 食材 -> WholeIngredientRecognition")
print(" 2. 菜品 -> DishClassification")
print(" 3. 半加工 -> ProcessedIngredientRecognition")
print("="*80 + "\n")
for i, ds in enumerate(datasources, 1):
print(f"\n{'='*80}")
print(f"[{i}/3] 开始下载: {ds}")
print(f"{'='*80}\n")
batch_download_all(
datasource=ds,
output_dir=args.output,
max_images=args.max,
skip_existing=not args.no_skip
)
if i < len(datasources):
print(f"\n等待5秒后继续下一个数据源...")
import time
time.sleep(5)
print("\n" + "="*80)
print("✓ 所有数据源下载完成!")
print("="*80)
else:
batch_download_all(
datasource=args.datasource,
output_dir=args.output,
max_images=args.max,
skip_existing=not args.no_skip
)