"""YOLO 数据集划分:可创建输出文件夹,自定义 train/val/test 百分比(带界面)。""" from __future__ import annotations import argparse import random import shutil import tkinter as tk from pathlib import Path from tkinter import filedialog, messagebox, ttk import yaml ROOT = Path(__file__).resolve().parent IMAGE_EXTS = {".jpg", ".jpeg", ".png", ".bmp", ".webp", ".tif", ".tiff"} def iter_images(folder: Path) -> list[Path]: if not folder.is_dir(): return [] return sorted(p for p in folder.iterdir() if p.is_file() and p.suffix.lower() in IMAGE_EXTS) def find_label_for(img: Path, src: Path, split_hint: str | None = None) -> Path | None: """按 stem 在常见 labels 位置查找同名 txt。""" candidates: list[Path] = [] if split_hint: candidates.append(src / "labels" / split_hint / f"{img.stem}.txt") candidates.extend( [ src / "labels" / "all" / f"{img.stem}.txt", src / "labels" / f"{img.stem}.txt", img.with_suffix(".txt"), img.parent.parent / "labels" / img.parent.name / f"{img.stem}.txt", img.parent.parent / "labels" / f"{img.stem}.txt", ] ) for cand in candidates: if cand.exists(): return cand return None def find_image_label_pairs(src: Path) -> list[tuple[Path, Path | None]]: """收集源目录全部样本。 支持: - images/{all|train|val|test} + labels/... - images/ 扁平目录 - 根目录直接放图片 若已有 train/val,会合并后再重新划分。 """ pairs_map: dict[str, tuple[Path, Path | None]] = {} split_dirs = [] for name in ("all", "train", "val", "test"): d = src / "images" / name if d.is_dir() and any(d.iterdir()): split_dirs.append((name if name != "all" else None, d)) if split_dirs: for hint, img_dir in split_dirs: for img in iter_images(img_dir): pairs_map[img.stem] = (img, find_label_for(img, src, hint)) else: flat_candidates = [src / "images", src] img_dir = next((p for p in flat_candidates if iter_images(p)), None) if img_dir is None: raise FileNotFoundError(f"找不到图片目录,请检查: {src}") for img in iter_images(img_dir): pairs_map[img.stem] = (img, find_label_for(img, src, None)) pairs = [pairs_map[k] for k in sorted(pairs_map)] if not pairs: raise FileNotFoundError(f"在 {src} 下没有找到图片") return pairs def validate_pairs(pairs: list[tuple[Path, Path | None]]) -> list[str]: warnings: list[str] = [] missing = [img.name for img, lb in pairs if lb is None] if missing: warnings.append(f"有 {len(missing)} 张图缺少标签,例如: {', '.join(missing[:3])}") bad = 0 for img, lb in pairs: if lb is None: continue text = lb.read_text(encoding="utf-8").strip() if not text: continue for line in text.splitlines(): parts = line.split() if len(parts) != 5: bad += 1 continue try: vals = list(map(float, parts[1:])) except ValueError: bad += 1 continue if any(v < 0 or v > 1 for v in vals): bad += 1 if bad: warnings.append(f"有 {bad} 行标签格式异常(应为: class x y w h,且坐标 0~1)") return warnings def split_indices(n: int, train_r: float, val_r: float, test_r: float, seed: int) -> dict[str, list[int]]: total = train_r + val_r + test_r if abs(total - 1.0) > 1e-6: raise ValueError(f"train+val+test 必须等于 100%,当前为 {total * 100:.1f}%") if n < 2: raise ValueError("样本太少,至少需要 2 张图") if test_r > 0 and n < 3: raise ValueError("划分测试集至少需要 3 张图") idxs = list(range(n)) random.Random(seed).shuffle(idxs) n_test = int(round(n * test_r)) if test_r > 0 else 0 n_val = int(round(n * val_r)) n_train = n - n_val - n_test if n_train < 1: raise ValueError("训练集为空,请调大训练比例") if n_val < 1: raise ValueError("验证集为空,请调大验证比例") if test_r > 0 and n_test < 1: raise ValueError("测试集为空,请调大测试比例") return { "train": idxs[:n_train], "val": idxs[n_train : n_train + n_val], "test": idxs[n_train + n_val :] if n_test else [], } def transfer(src: Path, dst: Path, move: bool) -> None: dst.parent.mkdir(parents=True, exist_ok=True) if move: shutil.move(str(src), str(dst)) else: shutil.copy2(src, dst) def write_data_yaml(yaml_path: Path, dataset_root: Path, has_test: bool, names: dict[int, str]) -> None: try: rel = dataset_root.resolve().relative_to(ROOT.resolve()).as_posix() except ValueError: rel = dataset_root.resolve().as_posix() data = { "path": rel, "train": "images/train", "val": "images/val", "test": "images/test" if has_test else None, "names": names, } if data["test"] is None: data.pop("test") yaml_path.parent.mkdir(parents=True, exist_ok=True) with yaml_path.open("w", encoding="utf-8") as f: f.write("# 由 split_dataset.py 自动生成\n") f.write("# 训练: yolo train cfg=train.yaml\n\n") yaml.safe_dump(data, f, allow_unicode=True, sort_keys=False) def load_or_build_names(yaml_path: Path, names_arg: str) -> dict[int, str]: if names_arg.strip(): return {i: n.strip() for i, n in enumerate(names_arg.split(",")) if n.strip()} if yaml_path.exists(): with yaml_path.open(encoding="utf-8") as f: old = yaml.safe_load(f) or {} old_names = old.get("names") if isinstance(old_names, dict): return {int(k): str(v) for k, v in old_names.items()} if isinstance(old_names, list): return {i: str(v) for i, v in enumerate(old_names)} return {0: "object"} def run_split( src: Path, dst: Path, train_pct: float, val_pct: float, test_pct: float, seed: int, move: bool, names_arg: str, yaml_path: Path, ) -> str: if not src.exists(): raise FileNotFoundError(f"源目录不存在: {src}") dst.mkdir(parents=True, exist_ok=True) pairs = find_image_label_pairs(src) warnings = validate_pairs(pairs) missing = sum(1 for _, lb in pairs if lb is None) splits = split_indices(len(pairs), train_pct / 100.0, val_pct / 100.0, test_pct / 100.0, seed) same_dir = src.resolve() == dst.resolve() stage = dst / "_split_stage" if stage.exists(): shutil.rmtree(stage) stage_img = stage / "images" stage_lbl = stage / "labels" stage_img.mkdir(parents=True) stage_lbl.mkdir(parents=True) staged: list[tuple[Path, Path | None]] = [] for img, label in pairs: new_img = stage_img / img.name transfer(img, new_img, move=False) new_lbl = None if label is not None: new_lbl = stage_lbl / label.name transfer(label, new_lbl, move=False) staged.append((new_img, new_lbl)) if move: # 源文件已暂存后,再删除源侧原文件 for img, label in pairs: img.unlink(missing_ok=True) if label is not None: label.unlink(missing_ok=True) for split in ("train", "val", "test"): for kind in ("images", "labels"): out = dst / kind / split if out.exists(): shutil.rmtree(out) out.mkdir(parents=True, exist_ok=True) for split, idxs in splits.items(): if not idxs: continue img_out = dst / "images" / split lbl_out = dst / "labels" / split for i in idxs: img, label = staged[i] transfer(img, img_out / img.name, move=True) if label is not None: transfer(label, lbl_out / label.name, move=True) else: (lbl_out / f"{img.stem}.txt").write_text("", encoding="utf-8") shutil.rmtree(stage, ignore_errors=True) names = load_or_build_names(yaml_path, names_arg) write_data_yaml(yaml_path, dst, has_test=bool(splits["test"]), names=names) lines = [ f"源目录: {src}", f"输出目录: {dst}" + ("(原地重划分)" if same_dir else ""), f"总样本: {len(pairs)} 无标签: {missing}", ] for split, idxs in splits.items(): if idxs: lines.append(f" {split}: {len(idxs)}") for w in warnings: lines.append(f"警告: {w}") lines.append(f"已生成: {yaml_path}") return "\n".join(lines) class SplitApp(tk.Tk): def __init__(self) -> None: super().__init__() self.title("YOLO 数据集划分") self.geometry("760x520") self.minsize(700, 480) self.var_src = tk.StringVar(value=str(ROOT / "datasets" / "phonedata")) self.var_dst = tk.StringVar(value=str(ROOT / "datasets" / "phonedata")) self.var_yaml = tk.StringVar(value=str(ROOT / "data.yaml")) self.var_names = tk.StringVar(value="phone") self.var_train = tk.DoubleVar(value=80) self.var_val = tk.DoubleVar(value=20) self.var_test = tk.DoubleVar(value=0) self.var_use_test = tk.BooleanVar(value=False) self.var_seed = tk.IntVar(value=42) self.var_move = tk.BooleanVar(value=False) self.var_sum = tk.StringVar(value="合计: 100%") self._build() self.on_test_toggle() self.refresh_sum() def _build(self) -> None: pad = {"padx": 8, "pady": 6} frm = ttk.Frame(self, padding=12) frm.pack(fill="both", expand=True) frm.columnconfigure(1, weight=1) ttk.Label(frm, text="源数据文件夹").grid(row=0, column=0, sticky="w", **pad) ttk.Entry(frm, textvariable=self.var_src).grid(row=0, column=1, sticky="ew", **pad) ttk.Button(frm, text="浏览", command=self.browse_src).grid(row=0, column=2, **pad) ttk.Label(frm, text="输出文件夹").grid(row=1, column=0, sticky="w", **pad) ttk.Entry(frm, textvariable=self.var_dst).grid(row=1, column=1, sticky="ew", **pad) btns = ttk.Frame(frm) btns.grid(row=1, column=2, **pad) ttk.Button(btns, text="浏览", command=self.browse_dst).pack(side="left") ttk.Button(btns, text="新建", command=self.create_dst).pack(side="left", padx=(6, 0)) ttk.Label(frm, text="data.yaml").grid(row=2, column=0, sticky="w", **pad) ttk.Entry(frm, textvariable=self.var_yaml).grid(row=2, column=1, sticky="ew", **pad) ttk.Button(frm, text="浏览", command=self.browse_yaml).grid(row=2, column=2, **pad) ttk.Label(frm, text="类别名").grid(row=3, column=0, sticky="w", **pad) ttk.Entry(frm, textvariable=self.var_names).grid(row=3, column=1, sticky="ew", **pad) ttk.Label(frm, text="逗号分隔,如 phone,cable").grid(row=3, column=2, sticky="w", **pad) ratio = ttk.LabelFrame(frm, text="划分百分比(合计必须 100%)", padding=10) ratio.grid(row=4, column=0, columnspan=3, sticky="ew", **pad) for c in range(6): ratio.columnconfigure(c, weight=1) ttk.Label(ratio, text="训练 %").grid(row=0, column=0, sticky="w") train_entry = ttk.Entry(ratio, textvariable=self.var_train, width=8) train_entry.grid(row=0, column=1, sticky="w") ttk.Label(ratio, text="验证 %").grid(row=0, column=2, sticky="w") val_entry = ttk.Entry(ratio, textvariable=self.var_val, width=8) val_entry.grid(row=0, column=3, sticky="w") ttk.Checkbutton(ratio, text="启用测试集", variable=self.var_use_test, command=self.on_test_toggle).grid( row=0, column=4, sticky="w" ) self.test_entry = ttk.Entry(ratio, textvariable=self.var_test, width=8) self.test_entry.grid(row=0, column=5, sticky="w") ttk.Label(ratio, textvariable=self.var_sum).grid(row=1, column=0, columnspan=3, sticky="w", pady=(8, 0)) ttk.Label(ratio, text="快捷:").grid(row=1, column=3, sticky="e", pady=(8, 0)) ttk.Button(ratio, text="80/20", command=lambda: self.set_ratio(80, 20, 0)).grid(row=1, column=4, sticky="w", pady=(8, 0)) ttk.Button(ratio, text="70/20/10", command=lambda: self.set_ratio(70, 20, 10)).grid( row=1, column=5, sticky="w", pady=(8, 0) ) for var in (self.var_train, self.var_val, self.var_test): var.trace_add("write", lambda *_: self.refresh_sum()) opts = ttk.Frame(frm) opts.grid(row=5, column=0, columnspan=3, sticky="ew", **pad) ttk.Label(opts, text="随机种子").pack(side="left") ttk.Entry(opts, textvariable=self.var_seed, width=8).pack(side="left", padx=(6, 16)) ttk.Checkbutton(opts, text="移动文件(默认复制)", variable=self.var_move).pack(side="left") ttk.Button(frm, text="扫描源目录", command=self.scan_src).grid(row=6, column=0, columnspan=1, sticky="ew", pady=10) ttk.Button(frm, text="开始划分", command=self.run).grid(row=6, column=1, columnspan=2, sticky="ew", pady=10) self.log = tk.Text(frm, height=12, wrap="word") self.log.grid(row=7, column=0, columnspan=3, sticky="nsew") frm.rowconfigure(7, weight=1) self.log.insert("end", "说明:\n") self.log.insert("end", "1. 源目录可为未划分数据,或已有 train/val(会合并后重新划分)\n") self.log.insert("end", "2. 输出目录可新建;结构: images/train|val|test + labels/...\n") self.log.insert("end", "3. 百分比自定义,启用测试集时 train+val+test=100\n") self.log.insert("end", "4. 当前 phonedata 格式正确时可直接点「扫描源目录」确认数量\n") def browse_src(self) -> None: path = filedialog.askdirectory(initialdir=self.var_src.get() or str(ROOT)) if path: self.var_src.set(path) def browse_dst(self) -> None: path = filedialog.askdirectory(initialdir=self.var_dst.get() or str(ROOT)) if path: self.var_dst.set(path) def create_dst(self) -> None: parent = filedialog.askdirectory(title="选择新建文件夹的父目录", initialdir=str(ROOT / "datasets")) if not parent: return dialog = tk.Toplevel(self) dialog.title("新建输出文件夹") dialog.geometry("420x140") dialog.transient(self) dialog.grab_set() name_var = tk.StringVar(value="split") ttk.Label(dialog, text=f"父目录: {parent}").pack(anchor="w", padx=12, pady=(12, 4)) row = ttk.Frame(dialog) row.pack(fill="x", padx=12) ttk.Label(row, text="文件夹名").pack(side="left") ttk.Entry(row, textvariable=name_var).pack(side="left", fill="x", expand=True, padx=8) def ok() -> None: name = name_var.get().strip() if not name: messagebox.showwarning("新建", "请输入文件夹名", parent=dialog) return path = Path(parent) / name path.mkdir(parents=True, exist_ok=True) self.var_dst.set(str(path)) dialog.destroy() ttk.Button(dialog, text="创建", command=ok).pack(pady=12) def browse_yaml(self) -> None: path = filedialog.asksaveasfilename( initialdir=str(ROOT), initialfile="data.yaml", defaultextension=".yaml", filetypes=[("YAML", "*.yaml;*.yml"), ("All", "*.*")], ) if path: self.var_yaml.set(path) def set_ratio(self, train: float, val: float, test: float) -> None: self.var_use_test.set(test > 0) self.on_test_toggle() self.var_train.set(train) self.var_val.set(val) self.var_test.set(test) self.refresh_sum() def on_test_toggle(self) -> None: if self.var_use_test.get(): self.test_entry.state(["!disabled"]) if float(self.var_test.get() or 0) <= 0: self.var_test.set(10) self.var_train.set(70) self.var_val.set(20) else: self.var_test.set(0) self.test_entry.state(["disabled"]) self.refresh_sum() def refresh_sum(self) -> None: try: train = float(self.var_train.get()) val = float(self.var_val.get()) test = float(self.var_test.get()) if self.var_use_test.get() else 0.0 total = train + val + test ok = abs(total - 100.0) < 0.05 self.var_sum.set(f"合计: {total:.1f}% {'✓' if ok else '(必须等于 100%)'}") except (tk.TclError, ValueError, TypeError): self.var_sum.set("合计: --") def scan_src(self) -> None: try: pairs = find_image_label_pairs(Path(self.var_src.get())) warnings = validate_pairs(pairs) missing = sum(1 for _, lb in pairs if lb is None) msg = f"扫描完成: 共 {len(pairs)} 张图,缺标签 {missing}\n" for w in warnings: msg += f"警告: {w}\n" if not warnings and missing == 0: msg += "图文配对与标签格式检查通过。\n" except Exception as exc: messagebox.showerror("扫描失败", str(exc)) self.log.insert("end", f"\n扫描错误: {exc}\n") return self.log.insert("end", f"\n{msg}") self.log.see("end") messagebox.showinfo("扫描结果", msg) def run(self) -> None: try: train = float(self.var_train.get()) val = float(self.var_val.get()) test = float(self.var_test.get()) if self.var_use_test.get() else 0.0 if abs(train + val + test - 100.0) > 0.05: raise ValueError("百分比合计必须等于 100") msg = run_split( src=Path(self.var_src.get()), dst=Path(self.var_dst.get()), train_pct=train, val_pct=val, test_pct=test, seed=int(self.var_seed.get()), move=self.var_move.get(), names_arg=self.var_names.get(), yaml_path=Path(self.var_yaml.get()), ) except Exception as exc: messagebox.showerror("划分失败", str(exc)) self.log.insert("end", f"\n错误: {exc}\n") return self.log.insert("end", f"\n{msg}\n") self.log.see("end") messagebox.showinfo("完成", msg) def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description="YOLO 数据集划分") parser.add_argument("--cli", action="store_true", help="命令行模式(默认打开界面)") parser.add_argument("--src", type=Path, default=ROOT / "datasets" / "phonedata") parser.add_argument("--dst", type=Path, default=ROOT / "datasets" / "phonedata") parser.add_argument("--train", type=float, default=80, help="训练集百分比,如 80") parser.add_argument("--val", type=float, default=20, help="验证集百分比,如 20") parser.add_argument("--test", type=float, default=0, help="测试集百分比,如 10") parser.add_argument("--seed", type=int, default=42) parser.add_argument("--move", action="store_true") parser.add_argument("--names", type=str, default="phone") parser.add_argument("--yaml", type=Path, default=ROOT / "data.yaml") return parser.parse_args() def main() -> None: args = parse_args() if args.cli: msg = run_split( src=args.src if args.src.is_absolute() else ROOT / args.src, dst=args.dst if args.dst.is_absolute() else ROOT / args.dst, train_pct=args.train, val_pct=args.val, test_pct=args.test, seed=args.seed, move=args.move, names_arg=args.names, yaml_path=args.yaml if args.yaml.is_absolute() else ROOT / args.yaml, ) print(msg) return SplitApp().mainloop() if __name__ == "__main__": main()