goat / Scripts /split_dataset.py
LightChuan's picture
Upload folder using huggingface_hub
6a5bb7e verified
Raw
History Blame Contribute Delete
2.27 kB
"""将已转换的YOLO标注和对应图片划分为train/val集"""
import shutil
import random
import argparse
from pathlib import Path
CAMERAS = ["EastLeft", "EastRight", "WestRight", "Westleft"]
PERIODS = ["Day", "Night"]
def main():
parser = argparse.ArgumentParser()
parser.add_argument("labeled_images_root", help="Labeled_images目录")
parser.add_argument("yolo_labels_root", help="yolo_labels_tmp目录")
parser.add_argument("out_root", help="Detection_dataset目录")
parser.add_argument("--val-ratio", type=float, default=0.1)
parser.add_argument("--seed", type=int, default=42)
args = parser.parse_args()
random.seed(args.seed)
labeled_root = Path(args.labeled_images_root)
labels_root = Path(args.yolo_labels_root)
out_root = Path(args.out_root)
train_img = out_root / "images" / "train"
val_img = out_root / "images" / "val"
train_lbl = out_root / "labels" / "train"
val_lbl = out_root / "labels" / "val"
for d in [train_img, val_img, train_lbl, val_lbl]:
d.mkdir(parents=True, exist_ok=True)
stats = {"train": 0, "val": 0}
for cam in CAMERAS:
for period in PERIODS:
img_dir = labeled_root / cam / period
lbl_dir = labels_root / cam / f"{period}Label"
if not img_dir.exists() or not lbl_dir.exists():
print(f"跳过 {cam}/{period}(目录不存在)")
continue
imgs = sorted(img_dir.glob("*.jpg"))
random.shuffle(imgs)
n_val = max(1, int(len(imgs) * args.val_ratio))
val_set = set(img.stem for img in imgs[:n_val])
for img_path in imgs:
lbl_path = lbl_dir / (img_path.stem + ".txt")
if not lbl_path.exists():
print(f"警告:找不到标注 {lbl_path}")
continue
split = "val" if img_path.stem in val_set else "train"
shutil.copy2(img_path, out_root / "images" / split / img_path.name)
shutil.copy2(lbl_path, out_root / "labels" / split / (img_path.stem + ".txt"))
stats[split] += 1
print(f"划分完成: train={stats['train']}, val={stats['val']}")
if __name__ == "__main__":
main()