Guest User

Untitled

a guest
May 17th, 2026
124
0
Never
Not a member of Pastebin yet? Sign Up, it unlocks many cool features!
text 38.34 KB | None | 0 0
  1. """
  2. Beat Scripter GUI
  3. =================
  4.  
  5. A Tkinter wrapper around the BeatScripter workflow from the original Jupyter notebook.
  6.  
  7. Workflow (one tab per step):
  8. 1. Setup - pick the video, choose start times, frames-per-section, model.
  9. 2. Box Editor - draw or type the bounding box on a sample frame.
  10. 3. Annotate - mark each frame as beat / not-beat (Space + B shortcuts + buttons).
  11. 4. Train / Predict - fit the classifier, then predict beats across the video.
  12. 5. Export - write the .funscript file.
  13.  
  14. Dependencies:
  15. pip install opencv-python numpy pillow scikit-learn
  16.  
  17. Run:
  18. python beat_scripter_gui.py
  19. """
  20.  
  21. import json
  22. import os
  23. import queue
  24. import threading
  25. from pathlib import Path
  26.  
  27. import cv2
  28. import numpy as np
  29. from PIL import Image, ImageTk
  30.  
  31. import tkinter as tk
  32. from tkinter import ttk, filedialog, messagebox
  33.  
  34. from sklearn.ensemble import RandomForestClassifier
  35. from sklearn.neighbors import KNeighborsClassifier
  36. from sklearn.svm import SVC
  37.  
  38.  
  39. # ---------------------------------------------------------------------------
  40. # Core logic (re-packaged from the original notebook, GUI-free)
  41. # ---------------------------------------------------------------------------
  42.  
  43. class BeatScripter:
  44. """Headless version of the original BeatScripter class.
  45.  
  46. GUI-dependent methods (cv2.imshow, waitKey, beat_marker) have been removed —
  47. annotation happens in the Tkinter UI instead.
  48. """
  49.  
  50. MODELS = {
  51. "KNN (k=5)": lambda: KNeighborsClassifier(n_neighbors=5),
  52. "SVM (RBF)": lambda: SVC(),
  53. "Random Forest": lambda: RandomForestClassifier(),
  54. }
  55.  
  56. def __init__(self, file_path, output_path=None):
  57. self.file_path = file_path
  58. self.file_path_out = output_path or str(Path(file_path).with_suffix(".funscript"))
  59. self.box_top_corner = (0, 0)
  60. self.box_size = (50, 50)
  61. self.frames_list = [] # list of sections; each section is list of (i, frame, time_ms)
  62. self.beats_trained_list = [] # list of sections; each section is list of (i, beat_bool)
  63. self.model = KNeighborsClassifier(n_neighbors=5)
  64. self.predicted_beats = []
  65.  
  66. # -- model -------------------------------------------------------------
  67. def set_model(self, name):
  68. factory = self.MODELS.get(name)
  69. if factory is not None:
  70. self.model = factory()
  71.  
  72. # -- loading -----------------------------------------------------------
  73. def load_frames(self, start_time_in_msec, number_of_frames):
  74. if not os.path.isfile(self.file_path):
  75. raise FileNotFoundError(f"Video file not found: {self.file_path}")
  76. cap = cv2.VideoCapture(self.file_path)
  77. if not cap.isOpened():
  78. raise RuntimeError(f"Could not open video: {self.file_path}")
  79. if start_time_in_msec > 0:
  80. cap.set(cv2.CAP_PROP_POS_MSEC, float(start_time_in_msec))
  81. frames = []
  82. while len(frames) < number_of_frames:
  83. ok, frame = cap.read()
  84. if not ok:
  85. break
  86. # POS_FRAMES is index of the NEXT frame to read; subtract 1 for the frame just returned
  87. i = int(cap.get(cv2.CAP_PROP_POS_FRAMES)) - 1
  88. t = cap.get(cv2.CAP_PROP_POS_MSEC)
  89. frames.append((i, frame, t))
  90. cap.release()
  91. return frames
  92.  
  93. def load_frames_list(self, start_times_sec, number_of_frames=100, progress_cb=None):
  94. self.frames_list = []
  95. for idx, start_time in enumerate(start_times_sec):
  96. if progress_cb:
  97. progress_cb("section_start", idx, len(start_times_sec))
  98. frames = self.load_frames(start_time * 1000, number_of_frames)
  99. self.frames_list.append(frames)
  100. if progress_cb:
  101. progress_cb("section_done", idx, len(start_times_sec), len(frames))
  102.  
  103. # -- box --------------------------------------------------------------
  104. def get_sub_frame(self, frame):
  105. x, y = self.box_top_corner
  106. w, h = self.box_size
  107. return frame[y:y + h, x:x + w, :]
  108.  
  109. # -- training ---------------------------------------------------------
  110. def train(self):
  111. beats_train = []
  112. frames_train = []
  113. for b in self.beats_trained_list:
  114. beats_train += b
  115. for f in self.frames_list:
  116. frames_train += f
  117.  
  118. # Build a fast index from frame_id -> frame so we don't do O(n*m) lookups.
  119. frame_by_id = {i: frame for i, frame, _t in frames_train}
  120.  
  121. X, y = [], []
  122. for i, beat in beats_train:
  123. if i in frame_by_id:
  124. X.append(self.get_sub_frame(frame_by_id[i]).flatten())
  125. y.append(int(beat))
  126. if not X:
  127. raise ValueError("No training examples found — did you finish annotating?")
  128. self.model.fit(X, y)
  129. return len(X)
  130.  
  131. # -- prediction -------------------------------------------------------
  132. def predict_beats(self, number=None, progress_cb=None, cancel_flag=None):
  133. cap = cv2.VideoCapture(self.file_path)
  134. if not cap.isOpened():
  135. raise RuntimeError(f"Could not open video: {self.file_path}")
  136. total = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
  137. if number is None or number <= 0:
  138. number = total
  139. else:
  140. number = min(number, total)
  141.  
  142. beats = []
  143. for i in range(number):
  144. if cancel_flag is not None and cancel_flag.is_set():
  145. break
  146. ok, frame = cap.read()
  147. if not ok:
  148. break
  149. try:
  150. sub = self.get_sub_frame(frame)
  151. if sub.size == 0:
  152. continue
  153. X_pred = [sub.flatten()]
  154. pred = bool(self.model.predict(X_pred)[0] == 1)
  155. beats.append((pred, i, cap.get(cv2.CAP_PROP_POS_MSEC)))
  156. except Exception:
  157. break
  158. if progress_cb and (i % 25 == 0 or i == number - 1):
  159. progress_cb(i + 1, number)
  160. cap.release()
  161. self.predicted_beats = beats
  162. return len(beats)
  163.  
  164. # -- export -----------------------------------------------------------
  165. def export_funscript(self):
  166. actions = []
  167. pos = 0
  168. last_i = 0
  169. for beat, i, time in self.predicted_beats:
  170. if beat and (i - last_i) > 4:
  171. actions.append({"pos": pos, "at": int(round(time))})
  172. pos = 100 - pos
  173. last_i = i
  174. payload = {
  175. "version": "1.0",
  176. "inverted": False,
  177. "range": 90,
  178. "actions": actions,
  179. }
  180. with open(self.file_path_out, "w") as fh:
  181. json.dump(payload, fh)
  182. return len(actions)
  183.  
  184.  
  185. # ---------------------------------------------------------------------------
  186. # GUI
  187. # ---------------------------------------------------------------------------
  188.  
  189. DISPLAY_MAX_W = 960
  190. DISPLAY_MAX_H = 540
  191.  
  192.  
  193. def cv_to_photo(cv_frame, box=None, max_w=DISPLAY_MAX_W, max_h=DISPLAY_MAX_H):
  194. """Convert a BGR OpenCV frame to a Tk PhotoImage, optionally with a red box.
  195.  
  196. Returns (PhotoImage, scale, display_w, display_h).
  197. """
  198. h, w = cv_frame.shape[:2]
  199. scale = min(max_w / w, max_h / h, 1.0)
  200. new_w, new_h = max(1, int(w * scale)), max(1, int(h * scale))
  201. img = cv_frame.copy()
  202. if box is not None:
  203. x, y, bw, bh = box
  204. cv2.rectangle(img, (x, y), (x + bw, y + bh), (0, 0, 255), 2)
  205. img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
  206. img = cv2.resize(img, (new_w, new_h), interpolation=cv2.INTER_AREA)
  207. photo = ImageTk.PhotoImage(Image.fromarray(img))
  208. return photo, scale, new_w, new_h
  209.  
  210.  
  211. class BeatScripterGUI(tk.Tk):
  212. def __init__(self):
  213. super().__init__()
  214. self.title("Beat Scripter GUI")
  215. self.geometry("1180x820")
  216. self.minsize(900, 700)
  217.  
  218. self.scripter: BeatScripter | None = None
  219. self.event_queue = queue.Queue()
  220. self.cancel_flag = threading.Event()
  221. self.worker_thread = None
  222.  
  223. # Annotation runtime state
  224. self.ann_set_idx = 0
  225. self.ann_frame_idx = 0
  226. # Per-section annotation list: list of [(i, bool), ...] aligned with frames_list sections.
  227. self.annotations = []
  228.  
  229. # Box editor runtime state
  230. self.box_canvas_scale = 1.0
  231. self.box_drag_start = None
  232. self.box_drag_rect = None
  233.  
  234. self._build_ui()
  235. self._refresh_status("Ready. Pick a video to begin.")
  236.  
  237. # Poll for worker events
  238. self.after(100, self._poll_queue)
  239.  
  240. # Keyboard shortcuts for annotation — handled globally but gated by current tab
  241. self.bind_all("<space>", self._kb_space)
  242. self.bind_all("<KeyPress-b>", self._kb_b)
  243. self.bind_all("<KeyPress-B>", self._kb_b)
  244. self.bind_all("<Left>", self._kb_back)
  245.  
  246. # ---------------------------------------------------------------- UI
  247.  
  248. def _build_ui(self):
  249. self.notebook = ttk.Notebook(self)
  250. self.notebook.pack(fill="both", expand=True, padx=8, pady=8)
  251.  
  252. self.tab_setup = ttk.Frame(self.notebook)
  253. self.tab_box = ttk.Frame(self.notebook)
  254. self.tab_annotate = ttk.Frame(self.notebook)
  255. self.tab_predict = ttk.Frame(self.notebook)
  256. self.tab_export = ttk.Frame(self.notebook)
  257.  
  258. self.notebook.add(self.tab_setup, text="1. Setup")
  259. self.notebook.add(self.tab_box, text="2. Box Editor")
  260. self.notebook.add(self.tab_annotate, text="3. Annotate")
  261. self.notebook.add(self.tab_predict, text="4. Train / Predict")
  262. self.notebook.add(self.tab_export, text="5. Export")
  263.  
  264. self._build_setup_tab()
  265. self._build_box_tab()
  266. self._build_annotate_tab()
  267. self._build_predict_tab()
  268. self._build_export_tab()
  269.  
  270. self.status_var = tk.StringVar(value="")
  271. ttk.Separator(self, orient="horizontal").pack(side="bottom", fill="x")
  272. ttk.Label(self, textvariable=self.status_var, anchor="w", padding=(8, 4)).pack(
  273. side="bottom", fill="x"
  274. )
  275.  
  276. # -- Setup tab ---------------------------------------------------------
  277. def _build_setup_tab(self):
  278. f = self.tab_setup
  279. for i in range(2):
  280. f.columnconfigure(i, weight=1)
  281.  
  282. # Video path
  283. ttk.Label(f, text="Video file (.mp4):").grid(row=0, column=0, sticky="w", padx=8, pady=(12, 2))
  284. row1 = ttk.Frame(f)
  285. row1.grid(row=1, column=0, columnspan=2, sticky="ew", padx=8)
  286. row1.columnconfigure(0, weight=1)
  287. self.video_path_var = tk.StringVar()
  288. ttk.Entry(row1, textvariable=self.video_path_var).grid(row=0, column=0, sticky="ew")
  289. ttk.Button(row1, text="Browse...", command=self._pick_video).grid(row=0, column=1, padx=(6, 0))
  290.  
  291. # Output path
  292. ttk.Label(f, text="Output .funscript path:").grid(row=2, column=0, sticky="w", padx=8, pady=(12, 2))
  293. row2 = ttk.Frame(f)
  294. row2.grid(row=3, column=0, columnspan=2, sticky="ew", padx=8)
  295. row2.columnconfigure(0, weight=1)
  296. self.output_path_var = tk.StringVar()
  297. ttk.Entry(row2, textvariable=self.output_path_var).grid(row=0, column=0, sticky="ew")
  298. ttk.Button(row2, text="Save as...", command=self._pick_output).grid(row=0, column=1, padx=(6, 0))
  299.  
  300. # Start times
  301. ttk.Label(f, text="Start times (seconds, comma-separated):").grid(
  302. row=4, column=0, sticky="w", padx=8, pady=(12, 2)
  303. )
  304. self.start_times_var = tk.StringVar(value="0, 90, 300, 540")
  305. ttk.Entry(f, textvariable=self.start_times_var).grid(
  306. row=5, column=0, columnspan=2, sticky="ew", padx=8
  307. )
  308. ttk.Label(
  309. f,
  310. text="Tip: 0, 1.5*60, 5*60, 9*60 → 0, 90, 300, 540. Decimals allowed.",
  311. foreground="#666",
  312. ).grid(row=6, column=0, columnspan=2, sticky="w", padx=8)
  313.  
  314. # Frames per section
  315. ttk.Label(f, text="Frames per section:").grid(row=7, column=0, sticky="w", padx=8, pady=(12, 2))
  316. self.frames_per_section_var = tk.IntVar(value=100)
  317. ttk.Spinbox(f, from_=10, to=2000, increment=10, textvariable=self.frames_per_section_var, width=10).grid(
  318. row=8, column=0, sticky="w", padx=8
  319. )
  320.  
  321. # Model picker
  322. ttk.Label(f, text="Classifier:").grid(row=7, column=1, sticky="w", padx=8, pady=(12, 2))
  323. self.model_var = tk.StringVar(value="KNN (k=5)")
  324. ttk.Combobox(
  325. f,
  326. textvariable=self.model_var,
  327. values=list(BeatScripter.MODELS.keys()),
  328. state="readonly",
  329. width=20,
  330. ).grid(row=8, column=1, sticky="w", padx=8)
  331.  
  332. # Action row
  333. action_row = ttk.Frame(f)
  334. action_row.grid(row=9, column=0, columnspan=2, sticky="ew", padx=8, pady=(20, 4))
  335. ttk.Button(action_row, text="Load Frames", command=self._on_load_frames).pack(side="left")
  336. self.load_progress = ttk.Progressbar(action_row, mode="determinate", length=300)
  337. self.load_progress.pack(side="left", padx=12)
  338. self.load_info_var = tk.StringVar()
  339. ttk.Label(action_row, textvariable=self.load_info_var).pack(side="left")
  340.  
  341. # -- Box tab -----------------------------------------------------------
  342. def _build_box_tab(self):
  343. f = self.tab_box
  344. f.columnconfigure(0, weight=1)
  345. f.rowconfigure(1, weight=1)
  346.  
  347. # Controls row
  348. ctrl = ttk.Frame(f)
  349. ctrl.grid(row=0, column=0, sticky="ew", padx=8, pady=8)
  350.  
  351. ttk.Label(ctrl, text="Section:").pack(side="left")
  352. self.box_set_var = tk.StringVar()
  353. self.box_set_combo = ttk.Combobox(
  354. ctrl, textvariable=self.box_set_var, state="readonly", width=22
  355. )
  356. self.box_set_combo.pack(side="left", padx=(4, 12))
  357. self.box_set_combo.bind("<<ComboboxSelected>>", lambda *_: self._render_box_frame())
  358.  
  359. ttk.Label(ctrl, text="Frame:").pack(side="left")
  360. self.box_frame_var = tk.IntVar(value=0)
  361. self.box_frame_scale = ttk.Scale(
  362. ctrl,
  363. from_=0,
  364. to=0,
  365. orient="horizontal",
  366. length=220,
  367. variable=self.box_frame_var,
  368. command=lambda *_: self._render_box_frame(),
  369. )
  370. self.box_frame_scale.pack(side="left", padx=(4, 12))
  371.  
  372. ttk.Label(ctrl, text="x:").pack(side="left")
  373. self.box_x_var = tk.IntVar(value=0)
  374. ttk.Spinbox(ctrl, from_=0, to=9999, textvariable=self.box_x_var, width=6).pack(side="left")
  375. ttk.Label(ctrl, text="y:").pack(side="left", padx=(8, 0))
  376. self.box_y_var = tk.IntVar(value=0)
  377. ttk.Spinbox(ctrl, from_=0, to=9999, textvariable=self.box_y_var, width=6).pack(side="left")
  378. ttk.Label(ctrl, text="w:").pack(side="left", padx=(8, 0))
  379. self.box_w_var = tk.IntVar(value=50)
  380. ttk.Spinbox(ctrl, from_=1, to=9999, textvariable=self.box_w_var, width=6).pack(side="left")
  381. ttk.Label(ctrl, text="h:").pack(side="left", padx=(8, 0))
  382. self.box_h_var = tk.IntVar(value=50)
  383. ttk.Spinbox(ctrl, from_=1, to=9999, textvariable=self.box_h_var, width=6).pack(side="left")
  384.  
  385. ttk.Button(ctrl, text="Apply", command=self._apply_box).pack(side="left", padx=(12, 0))
  386. ttk.Label(
  387. ctrl,
  388. text="(click + drag on the image to draw a new box)",
  389. foreground="#666",
  390. ).pack(side="left", padx=(8, 0))
  391.  
  392. # Canvas
  393. canvas_wrap = ttk.Frame(f, relief="sunken", borderwidth=1)
  394. canvas_wrap.grid(row=1, column=0, sticky="nsew", padx=8, pady=(0, 8))
  395. self.box_canvas = tk.Canvas(canvas_wrap, background="#222", highlightthickness=0)
  396. self.box_canvas.pack(fill="both", expand=True)
  397. self.box_canvas.bind("<ButtonPress-1>", self._box_canvas_press)
  398. self.box_canvas.bind("<B1-Motion>", self._box_canvas_drag)
  399. self.box_canvas.bind("<ButtonRelease-1>", self._box_canvas_release)
  400.  
  401. # -- Annotate tab ------------------------------------------------------
  402. def _build_annotate_tab(self):
  403. f = self.tab_annotate
  404. f.columnconfigure(0, weight=1)
  405. f.rowconfigure(1, weight=1)
  406.  
  407. info = ttk.Frame(f)
  408. info.grid(row=0, column=0, sticky="ew", padx=8, pady=8)
  409. self.ann_info_var = tk.StringVar(value="Load frames and set a box first.")
  410. ttk.Label(info, textvariable=self.ann_info_var, font=("TkDefaultFont", 11, "bold")).pack(side="left")
  411.  
  412. self.ann_progress = ttk.Progressbar(info, mode="determinate", length=280)
  413. self.ann_progress.pack(side="right")
  414.  
  415. canvas_wrap = ttk.Frame(f, relief="sunken", borderwidth=1)
  416. canvas_wrap.grid(row=1, column=0, sticky="nsew", padx=8, pady=(0, 8))
  417. self.ann_canvas = tk.Canvas(canvas_wrap, background="#222", highlightthickness=0)
  418. self.ann_canvas.pack(fill="both", expand=True)
  419.  
  420. btns = ttk.Frame(f)
  421. btns.grid(row=2, column=0, sticky="ew", padx=8, pady=(0, 12))
  422. ttk.Button(btns, text="◀ Back (←)", command=self._ann_back).pack(side="left")
  423. ttk.Button(btns, text="Not Beat (Space)", command=lambda: self._ann_mark(False)).pack(
  424. side="left", padx=(8, 0)
  425. )
  426. ttk.Button(btns, text="Beat (B)", command=lambda: self._ann_mark(True)).pack(side="left", padx=(8, 0))
  427. ttk.Button(btns, text="Skip Section", command=self._ann_skip_section).pack(side="left", padx=(8, 0))
  428. ttk.Label(
  429. btns,
  430. text="Space = Not Beat, B = Beat, ← = Back",
  431. foreground="#666",
  432. ).pack(side="right")
  433.  
  434. # -- Predict tab -------------------------------------------------------
  435. def _build_predict_tab(self):
  436. f = self.tab_predict
  437. for i in range(2):
  438. f.columnconfigure(i, weight=1)
  439.  
  440. train_box = ttk.LabelFrame(f, text="Train")
  441. train_box.grid(row=0, column=0, columnspan=2, sticky="ew", padx=8, pady=8)
  442. ttk.Button(train_box, text="Train Classifier", command=self._on_train).pack(side="left", padx=8, pady=8)
  443. self.train_status_var = tk.StringVar(value="Not trained.")
  444. ttk.Label(train_box, textvariable=self.train_status_var).pack(side="left", padx=(8, 8))
  445.  
  446. pred_box = ttk.LabelFrame(f, text="Predict")
  447. pred_box.grid(row=1, column=0, columnspan=2, sticky="ew", padx=8, pady=8)
  448.  
  449. ttk.Label(pred_box, text="Number of frames (blank = all):").grid(
  450. row=0, column=0, sticky="w", padx=8, pady=8
  451. )
  452. self.predict_n_var = tk.StringVar(value="15000")
  453. ttk.Entry(pred_box, textvariable=self.predict_n_var, width=10).grid(row=0, column=1, sticky="w", padx=8)
  454.  
  455. self.predict_btn = ttk.Button(pred_box, text="Predict Beats", command=self._on_predict)
  456. self.predict_btn.grid(row=1, column=0, sticky="w", padx=8, pady=(0, 8))
  457. self.cancel_predict_btn = ttk.Button(
  458. pred_box, text="Cancel", command=self._cancel_predict, state="disabled"
  459. )
  460. self.cancel_predict_btn.grid(row=1, column=1, sticky="w", padx=8, pady=(0, 8))
  461.  
  462. self.predict_progress = ttk.Progressbar(pred_box, mode="determinate", length=560)
  463. self.predict_progress.grid(row=2, column=0, columnspan=2, sticky="ew", padx=8, pady=(0, 4))
  464. self.predict_info_var = tk.StringVar()
  465. ttk.Label(pred_box, textvariable=self.predict_info_var).grid(
  466. row=3, column=0, columnspan=2, sticky="w", padx=8, pady=(0, 8)
  467. )
  468.  
  469. # -- Export tab --------------------------------------------------------
  470. def _build_export_tab(self):
  471. f = self.tab_export
  472. f.columnconfigure(0, weight=1)
  473.  
  474. ttk.Label(f, text="Output .funscript path:").grid(row=0, column=0, sticky="w", padx=8, pady=(12, 2))
  475. row = ttk.Frame(f)
  476. row.grid(row=1, column=0, sticky="ew", padx=8)
  477. row.columnconfigure(0, weight=1)
  478. # We just mirror the same variable used on the Setup tab
  479. ttk.Entry(row, textvariable=self.output_path_var).grid(row=0, column=0, sticky="ew")
  480. ttk.Button(row, text="Save as...", command=self._pick_output).grid(row=0, column=1, padx=(6, 0))
  481.  
  482. self.export_summary_var = tk.StringVar(value="Run prediction first.")
  483. ttk.Label(f, textvariable=self.export_summary_var).grid(
  484. row=2, column=0, sticky="w", padx=8, pady=(16, 4)
  485. )
  486.  
  487. btn_row = ttk.Frame(f)
  488. btn_row.grid(row=3, column=0, sticky="w", padx=8, pady=(4, 8))
  489. ttk.Button(btn_row, text="Export .funscript", command=self._on_export).pack(side="left")
  490. ttk.Button(btn_row, text="Open output folder", command=self._open_output_folder).pack(
  491. side="left", padx=(8, 0)
  492. )
  493.  
  494. # ---------------------------------------------------------------- Setup actions
  495.  
  496. def _pick_video(self):
  497. path = filedialog.askopenfilename(
  498. title="Select a video file",
  499. filetypes=[("MP4 video", "*.mp4"), ("All video", "*.mp4 *.mov *.mkv *.avi"), ("All files", "*.*")],
  500. )
  501. if not path:
  502. return
  503. self.video_path_var.set(path)
  504. # Default output to same folder with .funscript extension
  505. out = str(Path(path).with_suffix(".funscript"))
  506. self.output_path_var.set(out)
  507. self._refresh_status(f"Selected video: {path}")
  508.  
  509. def _pick_output(self):
  510. path = filedialog.asksaveasfilename(
  511. title="Save funscript as",
  512. defaultextension=".funscript",
  513. filetypes=[("Funscript", "*.funscript"), ("JSON", "*.json"), ("All files", "*.*")],
  514. )
  515. if path:
  516. self.output_path_var.set(path)
  517.  
  518. def _parse_start_times(self):
  519. raw = self.start_times_var.get()
  520. out = []
  521. for tok in raw.split(","):
  522. tok = tok.strip()
  523. if not tok:
  524. continue
  525. # Allow simple "a*b" expressions like 1.5*60
  526. try:
  527. if "*" in tok:
  528. parts = [float(p.strip()) for p in tok.split("*")]
  529. val = 1.0
  530. for p in parts:
  531. val *= p
  532. else:
  533. val = float(tok)
  534. out.append(val)
  535. except ValueError:
  536. raise ValueError(f"Could not parse start time '{tok}'")
  537. if not out:
  538. raise ValueError("At least one start time is required.")
  539. return out
  540.  
  541. def _on_load_frames(self):
  542. video = self.video_path_var.get().strip()
  543. if not video:
  544. messagebox.showerror("No video", "Pick a video file first.")
  545. return
  546. if not os.path.isfile(video):
  547. messagebox.showerror("Missing file", f"Video file not found:\n{video}")
  548. return
  549. try:
  550. start_times = self._parse_start_times()
  551. except ValueError as exc:
  552. messagebox.showerror("Bad start times", str(exc))
  553. return
  554. try:
  555. n_frames = int(self.frames_per_section_var.get())
  556. except Exception:
  557. messagebox.showerror("Bad frames count", "Frames per section must be an integer.")
  558. return
  559.  
  560. out = self.output_path_var.get().strip() or None
  561. self.scripter = BeatScripter(video, output_path=out)
  562. self.scripter.set_model(self.model_var.get())
  563.  
  564. self.load_progress["value"] = 0
  565. self.load_progress["maximum"] = len(start_times)
  566. self.load_info_var.set("Loading...")
  567. self._refresh_status("Loading frames in background...")
  568.  
  569. def worker():
  570. def cb(*args):
  571. self.event_queue.put(("load_progress",) + args)
  572. try:
  573. self.scripter.load_frames_list(start_times, n_frames, progress_cb=cb)
  574. self.event_queue.put(("load_done", None))
  575. except Exception as exc:
  576. self.event_queue.put(("load_error", str(exc)))
  577.  
  578. self.worker_thread = threading.Thread(target=worker, daemon=True)
  579. self.worker_thread.start()
  580.  
  581. def _on_load_done(self):
  582. if not self.scripter or not self.scripter.frames_list:
  583. return
  584. # Populate Box Editor section combo
  585. section_labels = [
  586. f"Section {idx + 1} ({len(s)} frames)" for idx, s in enumerate(self.scripter.frames_list)
  587. ]
  588. self.box_set_combo["values"] = section_labels
  589. if section_labels:
  590. self.box_set_combo.current(0)
  591. first_section_len = len(self.scripter.frames_list[0]) if self.scripter.frames_list else 0
  592. self.box_frame_scale.configure(from_=0, to=max(0, first_section_len - 1))
  593. self.box_frame_var.set(0)
  594.  
  595. # Initialize annotations storage to empty for each section
  596. self.annotations = [[] for _ in self.scripter.frames_list]
  597. self.ann_set_idx = 0
  598. self.ann_frame_idx = 0
  599.  
  600. # Render first frame in box editor
  601. self._render_box_frame()
  602.  
  603. total = sum(len(s) for s in self.scripter.frames_list)
  604. self.load_info_var.set(f"Loaded {total} frames in {len(self.scripter.frames_list)} section(s).")
  605. self._refresh_status("Frames loaded. Set a box, then go to Annotate.")
  606. self.notebook.select(self.tab_box)
  607.  
  608. # ---------------------------------------------------------------- Box editor
  609.  
  610. def _selected_box_section_index(self):
  611. sel = self.box_set_combo.current()
  612. return max(0, sel)
  613.  
  614. def _render_box_frame(self):
  615. if not self.scripter or not self.scripter.frames_list:
  616. return
  617. sec = self._selected_box_section_index()
  618. if sec >= len(self.scripter.frames_list):
  619. return
  620. section = self.scripter.frames_list[sec]
  621. if not section:
  622. return
  623. idx = int(self.box_frame_var.get())
  624. idx = max(0, min(idx, len(section) - 1))
  625. _, frame, _ = section[idx]
  626. # Update scale based on current spinbox sizes
  627. self.box_frame_scale.configure(to=len(section) - 1)
  628. box = (
  629. int(self.box_x_var.get()),
  630. int(self.box_y_var.get()),
  631. int(self.box_w_var.get()),
  632. int(self.box_h_var.get()),
  633. )
  634. photo, scale, w, h = cv_to_photo(frame, box=box)
  635. self.box_canvas.config(width=w, height=h)
  636. self.box_canvas.delete("all")
  637. self.box_canvas.create_image(0, 0, anchor="nw", image=photo, tags=("img",))
  638. self.box_canvas.image = photo # keep ref
  639. self.box_canvas_scale = scale
  640.  
  641. def _box_canvas_press(self, event):
  642. self.box_drag_start = (event.x, event.y)
  643. if self.box_drag_rect is not None:
  644. self.box_canvas.delete(self.box_drag_rect)
  645. self.box_drag_rect = self.box_canvas.create_rectangle(
  646. event.x, event.y, event.x, event.y, outline="yellow", width=2
  647. )
  648.  
  649. def _box_canvas_drag(self, event):
  650. if self.box_drag_start is None or self.box_drag_rect is None:
  651. return
  652. x0, y0 = self.box_drag_start
  653. self.box_canvas.coords(self.box_drag_rect, x0, y0, event.x, event.y)
  654.  
  655. def _box_canvas_release(self, event):
  656. if self.box_drag_start is None:
  657. return
  658. x0, y0 = self.box_drag_start
  659. x1, y1 = event.x, event.y
  660. self.box_drag_start = None
  661. scale = self.box_canvas_scale or 1.0
  662. # Convert display coords to real coords
  663. rx0, ry0 = int(min(x0, x1) / scale), int(min(y0, y1) / scale)
  664. rw, rh = max(1, int(abs(x1 - x0) / scale)), max(1, int(abs(y1 - y0) / scale))
  665. self.box_x_var.set(rx0)
  666. self.box_y_var.set(ry0)
  667. self.box_w_var.set(rw)
  668. self.box_h_var.set(rh)
  669. self._apply_box()
  670.  
  671. def _apply_box(self):
  672. if not self.scripter:
  673. return
  674. try:
  675. x = int(self.box_x_var.get())
  676. y = int(self.box_y_var.get())
  677. w = max(1, int(self.box_w_var.get()))
  678. h = max(1, int(self.box_h_var.get()))
  679. except Exception:
  680. messagebox.showerror("Bad box", "Box coordinates must be integers.")
  681. return
  682. self.scripter.box_top_corner = (x, y)
  683. self.scripter.box_size = (w, h)
  684. self._render_box_frame()
  685. self._refresh_status(f"Box set to top=({x},{y}), size=({w}x{h}).")
  686.  
  687. # ---------------------------------------------------------------- Annotate
  688.  
  689. def _on_tab_change(self, *_):
  690. pass
  691.  
  692. def _is_annotate_tab(self):
  693. try:
  694. return self.notebook.index(self.notebook.select()) == 2
  695. except Exception:
  696. return False
  697.  
  698. def _kb_space(self, _evt):
  699. if not self._is_annotate_tab():
  700. return
  701. if isinstance(self.focus_get(), (ttk.Entry, tk.Entry, ttk.Spinbox, ttk.Combobox)):
  702. return
  703. self._ann_mark(False)
  704.  
  705. def _kb_b(self, _evt):
  706. if not self._is_annotate_tab():
  707. return
  708. if isinstance(self.focus_get(), (ttk.Entry, tk.Entry, ttk.Spinbox, ttk.Combobox)):
  709. return
  710. self._ann_mark(True)
  711.  
  712. def _kb_back(self, _evt):
  713. if not self._is_annotate_tab():
  714. return
  715. if isinstance(self.focus_get(), (ttk.Entry, tk.Entry, ttk.Spinbox, ttk.Combobox)):
  716. return
  717. self._ann_back()
  718.  
  719. def _ann_total_progress(self):
  720. if not self.scripter:
  721. return 0, 0
  722. total = sum(len(s) for s in self.scripter.frames_list)
  723. done = sum(len(a) for a in self.annotations)
  724. return done, total
  725.  
  726. def _render_ann_frame(self):
  727. if not self.scripter or not self.scripter.frames_list:
  728. self.ann_info_var.set("Load frames first.")
  729. return
  730. if self.ann_set_idx >= len(self.scripter.frames_list):
  731. self.ann_info_var.set("All sections annotated. Go to Train / Predict.")
  732. self.ann_canvas.delete("all")
  733. return
  734. section = self.scripter.frames_list[self.ann_set_idx]
  735. if self.ann_frame_idx >= len(section):
  736. # Move to next section
  737. self.ann_set_idx += 1
  738. self.ann_frame_idx = 0
  739. self._render_ann_frame()
  740. return
  741. _, frame, _ = section[self.ann_frame_idx]
  742. box = (
  743. self.scripter.box_top_corner[0],
  744. self.scripter.box_top_corner[1],
  745. self.scripter.box_size[0],
  746. self.scripter.box_size[1],
  747. )
  748. photo, _, w, h = cv_to_photo(frame, box=box)
  749. self.ann_canvas.config(width=w, height=h)
  750. self.ann_canvas.delete("all")
  751. self.ann_canvas.create_image(0, 0, anchor="nw", image=photo)
  752. self.ann_canvas.image = photo
  753.  
  754. done, total = self._ann_total_progress()
  755. self.ann_progress["maximum"] = max(total, 1)
  756. self.ann_progress["value"] = done
  757. self.ann_info_var.set(
  758. f"Section {self.ann_set_idx + 1}/{len(self.scripter.frames_list)} · "
  759. f"Frame {self.ann_frame_idx + 1}/{len(section)} · "
  760. f"Annotated {done}/{total}"
  761. )
  762.  
  763. def _ann_mark(self, beat: bool):
  764. if not self.scripter or not self.scripter.frames_list:
  765. return
  766. if self.ann_set_idx >= len(self.scripter.frames_list):
  767. return
  768. section = self.scripter.frames_list[self.ann_set_idx]
  769. if self.ann_frame_idx >= len(section):
  770. return
  771. i, _, _ = section[self.ann_frame_idx]
  772. # Make sure we don't double-record after Back
  773. sect_ann = self.annotations[self.ann_set_idx]
  774. # Trim any future annotations (if user went back then re-marked)
  775. sect_ann[:] = sect_ann[: self.ann_frame_idx]
  776. sect_ann.append((i, beat))
  777. self.ann_frame_idx += 1
  778. self._render_ann_frame()
  779.  
  780. def _ann_back(self):
  781. if self.ann_frame_idx > 0:
  782. self.ann_frame_idx -= 1
  783. elif self.ann_set_idx > 0:
  784. self.ann_set_idx -= 1
  785. self.ann_frame_idx = len(self.scripter.frames_list[self.ann_set_idx]) - 1
  786. # Remove the annotation we're going back over so the next mark overwrites cleanly
  787. sect_ann = self.annotations[self.ann_set_idx]
  788. if len(sect_ann) > self.ann_frame_idx:
  789. del sect_ann[self.ann_frame_idx:]
  790. self._render_ann_frame()
  791.  
  792. def _ann_skip_section(self):
  793. if not self.scripter:
  794. return
  795. self.ann_set_idx += 1
  796. self.ann_frame_idx = 0
  797. self._render_ann_frame()
  798.  
  799. # ---------------------------------------------------------------- Predict
  800.  
  801. def _on_train(self):
  802. if not self.scripter:
  803. messagebox.showerror("No data", "Load frames first.")
  804. return
  805. # Copy current annotations to the scripter
  806. self.scripter.beats_trained_list = [list(a) for a in self.annotations]
  807. self.scripter.set_model(self.model_var.get())
  808. try:
  809. n = self.scripter.train()
  810. except Exception as exc:
  811. messagebox.showerror("Training failed", str(exc))
  812. self.train_status_var.set("Training failed.")
  813. return
  814. self.train_status_var.set(f"Trained on {n} annotated frames.")
  815. self._refresh_status(f"Model trained on {n} frames.")
  816.  
  817. def _on_predict(self):
  818. if not self.scripter:
  819. messagebox.showerror("No data", "Load frames and train first.")
  820. return
  821. raw = self.predict_n_var.get().strip()
  822. n = None
  823. if raw:
  824. try:
  825. n = int(raw)
  826. except ValueError:
  827. messagebox.showerror("Bad number", "Number of frames must be an integer (or empty).")
  828. return
  829. self.cancel_flag.clear()
  830. self.predict_progress["value"] = 0
  831. self.predict_progress["maximum"] = max(n or 1, 1)
  832. self.predict_info_var.set("Predicting...")
  833. self.predict_btn.configure(state="disabled")
  834. self.cancel_predict_btn.configure(state="normal")
  835.  
  836. def worker():
  837. def cb(done, total):
  838. self.event_queue.put(("predict_progress", done, total))
  839. try:
  840. count = self.scripter.predict_beats(
  841. number=n, progress_cb=cb, cancel_flag=self.cancel_flag
  842. )
  843. self.event_queue.put(("predict_done", count))
  844. except Exception as exc:
  845. self.event_queue.put(("predict_error", str(exc)))
  846.  
  847. self.worker_thread = threading.Thread(target=worker, daemon=True)
  848. self.worker_thread.start()
  849.  
  850. def _cancel_predict(self):
  851. self.cancel_flag.set()
  852. self._refresh_status("Cancelling prediction...")
  853.  
  854. # ---------------------------------------------------------------- Export
  855.  
  856. def _on_export(self):
  857. if not self.scripter:
  858. messagebox.showerror("No data", "Run prediction first.")
  859. return
  860. if not self.scripter.predicted_beats:
  861. messagebox.showerror("No predictions", "There are no predicted beats to export.")
  862. return
  863. out = self.output_path_var.get().strip()
  864. if not out:
  865. messagebox.showerror("No output path", "Pick an output .funscript path.")
  866. return
  867. self.scripter.file_path_out = out
  868. try:
  869. n = self.scripter.export_funscript()
  870. except Exception as exc:
  871. messagebox.showerror("Export failed", str(exc))
  872. return
  873. self.export_summary_var.set(f"Wrote {n} actions to {out}")
  874. self._refresh_status(f"Funscript exported: {out}")
  875. messagebox.showinfo("Exported", f"Wrote {n} actions to:\n{out}")
  876.  
  877. def _open_output_folder(self):
  878. out = self.output_path_var.get().strip()
  879. if not out:
  880. return
  881. folder = str(Path(out).parent)
  882. try:
  883. if os.name == "nt":
  884. os.startfile(folder) # type: ignore[attr-defined]
  885. elif os.name == "posix":
  886. import subprocess
  887. subprocess.Popen(["xdg-open", folder])
  888. except Exception as exc:
  889. messagebox.showerror("Open folder failed", str(exc))
  890.  
  891. # ---------------------------------------------------------------- Queue
  892.  
  893. def _poll_queue(self):
  894. try:
  895. while True:
  896. event = self.event_queue.get_nowait()
  897. self._handle_event(event)
  898. except queue.Empty:
  899. pass
  900. self.after(80, self._poll_queue)
  901.  
  902. def _handle_event(self, event):
  903. tag = event[0]
  904. if tag == "load_progress":
  905. kind = event[1]
  906. if kind == "section_done":
  907. idx = event[2]
  908. self.load_progress["value"] = idx + 1
  909. self.load_info_var.set(f"Loaded section {idx + 1}/{event[3]} ({event[4]} frames)")
  910. elif tag == "load_done":
  911. self._on_load_done()
  912. elif tag == "load_error":
  913. messagebox.showerror("Load failed", event[1])
  914. self._refresh_status("Load failed.")
  915. elif tag == "predict_progress":
  916. done, total = event[1], event[2]
  917. self.predict_progress["maximum"] = max(total, 1)
  918. self.predict_progress["value"] = done
  919. self.predict_info_var.set(f"{done}/{total} frames processed")
  920. elif tag == "predict_done":
  921. count = event[1]
  922. self.predict_info_var.set(f"Predicted on {count} frames.")
  923. self.predict_btn.configure(state="normal")
  924. self.cancel_predict_btn.configure(state="disabled")
  925. self.export_summary_var.set(
  926. f"{count} frames scored. Click Export .funscript to write the file."
  927. )
  928. self._refresh_status("Prediction complete. Ready to export.")
  929. elif tag == "predict_error":
  930. self.predict_btn.configure(state="normal")
  931. self.cancel_predict_btn.configure(state="disabled")
  932. self.predict_info_var.set("Prediction failed.")
  933. messagebox.showerror("Prediction failed", event[1])
  934.  
  935. def _refresh_status(self, msg):
  936. self.status_var.set(msg)
  937.  
  938.  
  939. def main():
  940. app = BeatScripterGUI()
  941. # Re-render annotation view when switching tabs so it picks up box changes
  942. app.notebook.bind(
  943. "<<NotebookTabChanged>>",
  944. lambda *_: app._render_ann_frame() if app._is_annotate_tab() else None,
  945. )
  946. app.mainloop()
  947.  
  948.  
  949. if __name__ == "__main__":
  950. main()
  951.  
Advertisement
Add Comment
Please, Sign In to add comment