File size: 5,504 Bytes
4dc0321
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
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()