diff --git a/data_management/batch_download_all.py b/data_management/batch_download_all.py deleted file mode 100644 index 21725f3..0000000 --- a/data_management/batch_download_all.py +++ /dev/null @@ -1,298 +0,0 @@ -""" -批量下载所有图片到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 - )