mouse-gesture-recognizer / INTERFERENCE_GUi.py
nadizik's picture
Upload INTERFERENCE_GUi.py
4dc0321 verified
Raw
History Blame Contribute Delete
5.5 kB
import os
import sys
import tkinter as tk
from tkinter import ttk
import numpy as np
from scipy.interpolate import interp1d
import torch
import torch.nn as nn
import torch.nn.functional as F
def get_base_dir():
if getattr(sys, 'frozen', False):
return getattr(sys, '_MEIPASS', os.path.dirname(sys.executable))
return os.path.dirname(os.path.abspath(__file__))
class TaranCore(nn.Module):
"""Gesture Classification MLP matching exact model.pth tensor dimensions (128 -> 64 -> 32 -> 8)."""
def __init__(self, in_d=128, hid1_d=64, hid2_d=32, out_d=8):
super().__init__()
self.net = nn.Sequential(
nn.Linear(in_d, hid1_d),
nn.LeakyReLU(0.1),
nn.Linear(hid1_d, hid2_d),
nn.LeakyReLU(0.1),
nn.Linear(hid2_d, out_d)
)
def forward(self, x, temp=0.45):
if x.shape[0] == 1:
self.eval()
logits = self.net(x)
return F.softmax(logits / temp, dim=-1)
def process_pattern(points, target_pts=64):
pts = np.array(points, dtype=np.float32)
if len(pts) < 5:
return np.zeros(target_pts * 2, dtype=np.float32)
pts -= np.mean(pts, axis=0)
norm = np.max(np.abs(pts))
if norm > 0:
pts /= norm
try:
dist = np.sqrt(np.sum(np.diff(pts, axis=0) ** 2, axis=1))
dist = np.concatenate(([0], np.cumsum(dist)))
d_norm = dist / dist[-1]
d_norm, u_idx = np.unique(d_norm, return_index=True)
pts = pts[u_idx]
fx = interp1d(d_norm, pts[:, 0], fill_value="extrapolate")
fy = interp1d(d_norm, pts[:, 1], fill_value="extrapolate")
steps = np.linspace(0, 1, target_pts)
return np.vstack((fx(steps), fy(steps))).T.flatten()
except Exception:
return np.zeros(target_pts * 2, dtype=np.float32)
class GestureTesterApp:
def __init__(self, root, model):
self.root = root
self.root.title("TaranCore Gesture Recognizer - Interactive Test")
self.root.geometry("600x550")
self.root.resizable(False, False)
self.model = model
self.points = []
# UI Components
self.label_title = ttk.Label(root, text="Draw a gesture using Mouse or Touchpad", font=("Arial", 12, "bold"))
self.label_title.pack(pady=10)
self.canvas = tk.Canvas(root, width=400, height=300, bg="white", relief="ridge", bd=2)
self.canvas.pack(pady=5)
self.canvas.bind("<ButtonPress-1>", self.on_press)
self.canvas.bind("<B1-Motion>", self.on_drag)
self.canvas.bind("<ButtonRelease-1>", self.on_release)
self.label_result = ttk.Label(root, text="Result: Draw something...", font=("Arial", 14, "bold"), foreground="blue")
self.label_result.pack(pady=10)
self.frame_probs = ttk.Frame(root)
# Исправлено: padx=20 вместо px=20
self.frame_probs.pack(pady=5, fill="x", padx=20)
self.prob_labels = []
for i in range(8):
lbl = ttk.Label(self.frame_probs, text=f"Class {i}: 0.0%", font=("Consolas", 9))
lbl.grid(row=i // 4, column=i % 4, padx=15, pady=2)
self.prob_labels.append(lbl)
self.btn_clear = ttk.Button(root, text="Clear Canvas", command=self.clear_canvas)
self.btn_clear.pack(pady=15)
def on_press(self, event):
self.clear_canvas()
self.points.append((event.x, event.y))
def on_drag(self, event):
if self.points:
x_prev, y_prev = self.points[-1]
self.canvas.create_line(x_prev, y_prev, event.x, event.y, fill="black", width=3, capstyle=tk.ROUND, smooth=True)
self.points.append((event.x, event.y))
def on_release(self, event):
if len(self.points) < 5:
self.label_result.config(text="Result: Gesture too short!", foreground="red")
return
vector = process_pattern(self.points)
input_tensor = torch.FloatTensor(vector).unsqueeze(0)
with torch.no_grad():
probs = self.model(input_tensor).numpy().flatten()
predicted_class = int(np.argmax(probs))
confidence = float(probs[predicted_class] * 100)
self.label_result.config(
text=f"Detected: Class {predicted_class} ({confidence:.1f}%)",
foreground="green"
)
for idx, prob in enumerate(probs):
self.prob_labels[idx].config(text=f"Class {idx}: {prob * 100:.1f}%")
def clear_canvas(self):
self.canvas.delete("all")
self.points.clear()
self.label_result.config(text="Result: Draw something...", foreground="blue")
def main():
base_dir = get_base_dir()
weights_path = os.path.join(base_dir, "model.pth")
if not os.path.exists(weights_path):
print(f"[ERROR] Weights file not found: {weights_path}")
return
model = TaranCore(in_d=128, hid1_d=64, hid2_d=32, out_d=8)
try:
state_dict = torch.load(weights_path, map_location="cpu")
model.load_state_dict(state_dict)
model.eval()
except Exception as e:
print(f"[ERROR] Failed to load model state: {e}")
return
root = tk.Tk()
app = GestureTesterApp(root, model)
root.mainloop()
if __name__ == "__main__":
import multiprocessing
multiprocessing.freeze_support()
main()