first
This commit is contained in:
@@ -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()
|
||||
Reference in New Issue
Block a user