Files
ros_flexiv/split_dataset.py
orisys 4ad53f4e97 first
2026-08-21 14:51:57 +08:00

520 lines
20 KiB
Python
Raw Permalink 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.
"""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()