diff --git a/bash/pull_swe_bench_images.py b/bash/pull_swe_bench_images.py index c60048c..daa3002 100755 --- a/bash/pull_swe_bench_images.py +++ b/bash/pull_swe_bench_images.py @@ -111,6 +111,25 @@ def load_swe_samples(dataset_id: str): return loader.load() +def load_multiple_datasets(dataset_keys: list[str]) -> list: + """Load samples from one or more dataset keys and tag each sample with its source.""" + all_samples = [] + for key in dataset_keys: + dataset_id = DATASET_IDS[key] + samples = load_swe_samples(dataset_id) + # Tag each sample so get_image_name can still work and state stays per-dataset. + for sample in samples: + if isinstance(sample, dict): + sample.setdefault('dataset_key', key) + else: + metadata = getattr(sample, 'metadata', None) + if metadata is not None: + metadata.setdefault('dataset_key', key) + all_samples.extend(samples) + print(f' {key}: {len(samples)} 个样本') + return all_samples + + def get_image_name(instance, namespace: str, arch: str) -> str: """根据 swebench 规则生成远程镜像名称。""" if isinstance(instance, dict): @@ -217,8 +236,8 @@ def main(): parser.add_argument( '--dataset', default='swe_bench_verified', - choices=list(DATASET_IDS.keys()), - help='SWE-bench dataset variant', + type=lambda s: [x.strip() for x in s.split(',') if x.strip()], + help='SWE-bench dataset variant(s), comma-separated, e.g. swe_bench_verified or swe_bench_verified,swe_bench_lite', ) parser.add_argument( '--max-workers', @@ -264,14 +283,19 @@ def main(): state['failed'] = [] print(f'🔄 重试 {len(failed)} 个之前失败的镜像') - dataset_id = DATASET_IDS[args.dataset] - samples = load_swe_samples(dataset_id) + # Validate dataset keys + for key in args.dataset: + if key not in DATASET_IDS: + print(f'ERROR: unknown dataset "{key}". Supported: {list(DATASET_IDS.keys())}') + sys.exit(1) + + samples = load_multiple_datasets(args.dataset) if args.dry_run > 0: samples = samples[:args.dry_run] print(f'⚠️ 测试模式:仅处理前 {args.dry_run} 个样本以验证镜像名和网络...') else: - print(f'✅ 成功加载 {len(samples)} 个样本。') + print(f'✅ 成功加载 {len(samples)} 个样本(数据集: {args.dataset})。') already_done = len(state['done']) + len(state['skipped']) print(f'📂 进度文件: {get_state_path(output_dir)}') diff --git a/myread.md b/myread.md index 42d2ffe..5ac7387 100644 --- a/myread.md +++ b/myread.md @@ -243,6 +243,11 @@ export PYTHONUNBUFFERED=1 python bash/pull_swe_bench_images.py \ --dataset swe_bench_verified \ --max-workers 1 + +# 也可以同时预拉多个数据集(逗号分隔) +python bash/pull_swe_bench_images.py \ + --dataset swe_bench_verified,swe_bench_lite \ + --max-workers 1 ``` 脚本特性: @@ -250,6 +255,7 @@ python bash/pull_swe_bench_images.py \ - 断点续传:状态文件 `output/.swe_bench_pull_state.json`,中断后再次运行会自动跳过已成功的镜像。 - 本地已有镜像自动跳过。 - 单 worker 模式稳定,减少限流概率;若网络非常好,可适当提高 `--max-workers`。 +- `--dataset` 支持逗号分隔的多个数据集。 镜像会存到 Docker `data-root` 配置的路径: