This commit is contained in:
orisys
2026-08-21 14:51:57 +08:00
commit 4ad53f4e97
2056 changed files with 6272 additions and 0 deletions
+519
View File
@@ -0,0 +1,519 @@
"""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()