Isk5434 commited on
Commit
e3dcc12
·
verified ·
1 Parent(s): 858985e

Upload app.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. app.py +1553 -0
app.py ADDED
@@ -0,0 +1,1553 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # !pip -q install gradio fastapi uvicorn scikit-learn pyngrok scikit-learn qrcode[pil]
2
+ # =========================
3
+ # Colab-ready single script
4
+ # - Runs FastAPI + Gradio mounted app on Colab
5
+ # - Uses sid in JSON payload (cookie may be unreliable in Colab/iframes)
6
+ # =========================
7
+
8
+ # (1) Install deps (Colab only)
9
+ import sys, os, threading, time
10
+ if "google.colab" in sys.modules:
11
+ pass
12
+ # !pip -q install gradio fastapi uvicorn scikit-learn
13
+
14
+ import threading
15
+ from dataclasses import dataclass, field
16
+ from typing import Dict, List, Tuple, Optional
17
+ import numpy as np
18
+
19
+ # Headless matplotlib
20
+ import matplotlib
21
+ matplotlib.use("Agg")
22
+ import matplotlib.pyplot as plt # unused ok
23
+
24
+ import gradio as gr
25
+ from sklearn.model_selection import train_test_split
26
+
27
+ from fastapi import FastAPI, Request
28
+ from fastapi.responses import JSONResponse
29
+
30
+
31
+ # =========================
32
+ # Config
33
+ # =========================
34
+ @dataclass
35
+ class PreprocConfig:
36
+ fs_target: float = 50.0
37
+ hp_alpha: float = 0.92
38
+ window_sec: float = 1.0
39
+ hop_sec: float = 0.2
40
+
41
+ CFG = PreprocConfig()
42
+ DEFAULT_LABELS = ["idle", "shake", "flip"]
43
+
44
+
45
+ # =========================
46
+ # Utilities
47
+ # =========================
48
+ class GravityHighPass:
49
+ """Estimate gravity via EMA and subtract it: a_dyn = a - g_est"""
50
+ def __init__(self, alpha=0.92):
51
+ self.alpha = float(alpha)
52
+ self.g = np.zeros(3, dtype=np.float32)
53
+ self.inited = False
54
+
55
+ def reset(self):
56
+ self.g[:] = 0
57
+ self.inited = False
58
+
59
+ def step(self, a_xyz: np.ndarray) -> np.ndarray:
60
+ a = a_xyz.astype(np.float32)
61
+ if not self.inited:
62
+ self.g = a.copy()
63
+ self.inited = True
64
+ self.g = self.alpha * self.g + (1 - self.alpha) * a
65
+ return a - self.g
66
+
67
+
68
+ def resample_linear(ts: np.ndarray, X: np.ndarray, fs_target: float):
69
+ """Resample irregular timestamps to uniform grid using linear interpolation."""
70
+ if len(ts) < 2:
71
+ return ts, X
72
+ t0, t1 = float(ts[0]), float(ts[-1])
73
+ dt = 1.0 / fs_target
74
+ t_new = np.arange(t0, t1, dt, dtype=np.float32)
75
+ if len(t_new) < 2:
76
+ return ts, X
77
+ X_new = np.zeros((len(t_new), X.shape[1]), dtype=np.float32)
78
+ for d in range(X.shape[1]):
79
+ X_new[:, d] = np.interp(t_new, ts, X[:, d])
80
+ return t_new, X_new
81
+
82
+
83
+ # =========================
84
+ # ESN Classifier (minimal)
85
+ # =========================
86
+ class ESNClassifier:
87
+ """ESN state features + ridge multi-class regression."""
88
+ def __init__(self, in_dim: int, res_size: int, spectral_radius: float, leak: float, ridge: float, seed: int = 0):
89
+ self.in_dim = in_dim
90
+ self.res_size = res_size
91
+ self.spectral_radius = float(spectral_radius)
92
+ self.leak = float(leak)
93
+ self.ridge = float(ridge)
94
+ self.seed = int(seed)
95
+
96
+ rng = np.random.default_rng(self.seed)
97
+ self.Win = (rng.uniform(-1, 1, size=(res_size, in_dim + 1)) * 0.5).astype(np.float32)
98
+
99
+ W = rng.uniform(-1, 1, size=(res_size, res_size)).astype(np.float32)
100
+ v = rng.normal(size=(res_size,)).astype(np.float32)
101
+ for _ in range(30):
102
+ v = W @ v
103
+ v = v / (np.linalg.norm(v) + 1e-9)
104
+ eig_approx = float(np.linalg.norm(W @ v) / (np.linalg.norm(v) + 1e-9))
105
+ W *= (self.spectral_radius / (eig_approx + 1e-9))
106
+ self.W = W
107
+
108
+ self.x = np.zeros((res_size,), dtype=np.float32)
109
+ self.Wout = None
110
+ self.class_names: List[str] = []
111
+
112
+ def reset(self):
113
+ self.x[:] = 0
114
+
115
+ def step(self, u: np.ndarray):
116
+ u = u.astype(np.float32)
117
+ aug = np.concatenate([np.array([1.0], np.float32), u], axis=0)
118
+ pre = self.W @ self.x + self.Win @ aug
119
+ x_new = np.tanh(pre)
120
+ self.x = (1 - self.leak) * self.x + self.leak * x_new
121
+ return self.x
122
+
123
+ def fit(self, X_feat: np.ndarray, y: np.ndarray, class_names: List[str]):
124
+ self.class_names = class_names
125
+ n, f = X_feat.shape
126
+ k = len(class_names)
127
+ Y = np.zeros((n, k), dtype=np.float32)
128
+ Y[np.arange(n), y] = 1.0
129
+
130
+ XtX = X_feat.T @ X_feat
131
+ I = np.eye(f, dtype=np.float32)
132
+ self.Wout = np.linalg.solve(XtX + self.ridge * I, X_feat.T @ Y).astype(np.float32)
133
+
134
+ def predict_proba(self, feat: np.ndarray):
135
+ logits = feat.astype(np.float32) @ self.Wout
136
+ m = float(np.max(logits))
137
+ ex = np.exp(logits - m)
138
+ return ex / (float(np.sum(ex)) + 1e-9)
139
+
140
+
141
+ def make_window_feature(esn: ESNClassifier, X_seq: np.ndarray, mode: str = "last"):
142
+ esn.reset()
143
+ states = []
144
+ for t in range(len(X_seq)):
145
+ st = esn.step(X_seq[t])
146
+ if mode == "mean":
147
+ states.append(st.copy())
148
+ if mode == "mean" and len(states) > 0:
149
+ s = np.mean(np.stack(states, axis=0), axis=0)
150
+ else:
151
+ s = esn.x.copy()
152
+
153
+ u_mean = X_seq.mean(axis=0)
154
+ u_std = X_seq.std(axis=0)
155
+ feat = np.concatenate([np.array([1.0], np.float32), u_mean, u_std, s], axis=0)
156
+ return feat
157
+
158
+
159
+ # =========================
160
+ # Per-session state
161
+ # =========================
162
+ @dataclass
163
+ class SessionState:
164
+ stream_t: List[float] = field(default_factory=list)
165
+ stream_a: List[List[float]] = field(default_factory=list)
166
+
167
+ collecting: bool = False
168
+ collect_label: str = ""
169
+ collect_tmp_t: List[float] = field(default_factory=list)
170
+ collect_tmp_a: List[List[float]] = field(default_factory=list)
171
+
172
+ data: Dict[str, List[Dict[str, np.ndarray]]] = field(default_factory=dict)
173
+
174
+ trained: bool = False
175
+ train_cfg: Dict = field(default_factory=dict)
176
+ pp_mean: Optional[np.ndarray] = None
177
+ pp_std: Optional[np.ndarray] = None
178
+ esn_model: Optional[ESNClassifier] = None
179
+
180
+ infer_running: bool = False
181
+ infer_last_label: str = ""
182
+ infer_last_conf: float = 0.0
183
+ infer_pred_log: List[Tuple[float, str, float]] = field(default_factory=list)
184
+
185
+ lock: threading.Lock = field(default_factory=threading.Lock)
186
+
187
+
188
+ SESS: Dict[str, SessionState] = {}
189
+ SESS_LOCK = threading.Lock()
190
+
191
+
192
+ def _get_sid_from_request(request: Optional[gr.Request]) -> str:
193
+ if request is None:
194
+ return "unknown"
195
+ try:
196
+ sid = request.cookies.get("sid", "") if request.cookies else ""
197
+ return sid or "unknown"
198
+ except Exception:
199
+ return "unknown"
200
+
201
+
202
+ def get_state(sid: str) -> SessionState:
203
+ with SESS_LOCK:
204
+ st = SESS.get(sid)
205
+ if st is None:
206
+ st = SessionState()
207
+ SESS[sid] = st
208
+ return st
209
+
210
+
211
+ def reset_state(st: SessionState):
212
+ st.stream_t = []
213
+ st.stream_a = []
214
+ st.collecting = False
215
+ st.collect_label = ""
216
+ st.collect_tmp_t = []
217
+ st.collect_tmp_a = []
218
+ st.data = {}
219
+
220
+ st.trained = False
221
+ st.train_cfg = {}
222
+ st.pp_mean = None
223
+ st.pp_std = None
224
+ st.esn_model = None
225
+
226
+ st.infer_running = False
227
+ st.infer_last_label = ""
228
+ st.infer_last_conf = 0.0
229
+ st.infer_pred_log = []
230
+
231
+
232
+ def counts_dict_for(st: SessionState):
233
+ c = {k: len(v) for k, v in st.data.items()}
234
+ c["TOTAL"] = int(sum(c.values()))
235
+ return c
236
+
237
+
238
+ def ui_status_for(st: SessionState):
239
+ if not st.trained:
240
+ return "<span style='font-size:18px;font-weight:700;color:#b00020'>MODEL: not trained</span>"
241
+ return (f"<span style='font-size:18px;font-weight:700;color:#0b6b0b'>MODEL: trained</span> "
242
+ f"<span style='font-size:12px;opacity:.85'>val_acc={st.train_cfg.get('val_acc',0):.3f}, "
243
+ f"classes={st.train_cfg.get('classes',[])}</span>")
244
+
245
+
246
+ def format_pred_log_md(st: SessionState, max_rows: int = 80):
247
+ if len(st.infer_pred_log) == 0:
248
+ return "(log empty)"
249
+ rows = st.infer_pred_log[-max_rows:]
250
+ md = ["| time(s) | label | conf |", "|---:|:---|---:|"]
251
+ for t, lab, conf in rows:
252
+ md.append(f"| {t:6.2f} | {lab} | {conf:.2f} |")
253
+ return "\n".join(md)
254
+
255
+
256
+ # =========================
257
+ # FastAPI endpoints
258
+ # - prefer cookie sid; if missing, use payload sid
259
+ # =========================
260
+ api = FastAPI()
261
+
262
+ def _sid_from_fastapi(request: Request, payload: dict) -> str:
263
+ try:
264
+ sid = request.cookies.get("sid", "") or ""
265
+ except Exception:
266
+ sid = ""
267
+ if sid:
268
+ return sid
269
+ sid2 = str(payload.get("sid", "") or "")
270
+ return sid2 if sid2 else "unknown"
271
+
272
+
273
+ @api.post("/api/ingest")
274
+ async def ingest(request: Request):
275
+ try:
276
+ obj = await request.json()
277
+ sid = _sid_from_fastapi(request, obj)
278
+ samples = obj.get("samples", [])
279
+
280
+ st = get_state(sid)
281
+ with st.lock:
282
+ for s in samples:
283
+ t = float(s.get("t", 0.0))
284
+ ax = float(s.get("ax", 0.0)); ay = float(s.get("ay", 0.0)); az = float(s.get("az", 0.0))
285
+ st.stream_t.append(t)
286
+ st.stream_a.append([ax, ay, az])
287
+ if st.collecting:
288
+ st.collect_tmp_t.append(t)
289
+ st.collect_tmp_a.append([ax, ay, az])
290
+
291
+ if len(st.stream_t) > 6000:
292
+ st.stream_t = st.stream_t[-6000:]
293
+ st.stream_a = st.stream_a[-6000:]
294
+
295
+ return JSONResponse({"ok": True, "n": len(samples), "sid": sid, "stream_len": len(st.stream_t)})
296
+ except Exception as e:
297
+ return JSONResponse({"ok": False, "error": str(e)}, status_code=400)
298
+
299
+
300
+ @api.post("/api/reset")
301
+ async def reset_endpoint(request: Request):
302
+ try:
303
+ obj = await request.json()
304
+ sid = _sid_from_fastapi(request, obj)
305
+ st = get_state(sid)
306
+ with st.lock:
307
+ reset_state(st)
308
+ return JSONResponse({"ok": True, "sid": sid})
309
+ except Exception as e:
310
+ return JSONResponse({"ok": False, "error": str(e)}, status_code=400)
311
+
312
+
313
+ # =========================
314
+ # Collect handlers
315
+ # =========================
316
+ def collect_start(label: str, request: gr.Request):
317
+ sid = _get_sid_from_request(request)
318
+ st = get_state(sid)
319
+ with st.lock:
320
+ st.collecting = True
321
+ st.collect_label = label
322
+ st.collect_tmp_t = []
323
+ st.collect_tmp_a = []
324
+ return f"収集中: {label}", gr.update(interactive=False), gr.update(interactive=True)
325
+
326
+ def collect_stop_and_save(request: gr.Request):
327
+ sid = _get_sid_from_request(request)
328
+ st = get_state(sid)
329
+
330
+ with st.lock:
331
+ st.collecting = False
332
+
333
+ if len(st.collect_tmp_t) < 10:
334
+ return ("収集停止(データが少なすぎるため未保存)",
335
+ gr.update(interactive=True), gr.update(interactive=False),
336
+ counts_dict_for(st))
337
+
338
+ ts = np.array(st.collect_tmp_t, dtype=np.float32)
339
+ A = np.array(st.collect_tmp_a, dtype=np.float32)
340
+ lab = st.collect_label
341
+ st.data.setdefault(lab, []).append({"t": ts, "a": A})
342
+
343
+ msg = f"保存: label={lab}, samples={len(ts)}, total={len(st.data[lab])}"
344
+ return (msg, gr.update(interactive=True), gr.update(interactive=False),
345
+ counts_dict_for(st))
346
+
347
+
348
+ # =========================
349
+ # Training
350
+ # =========================
351
+ def build_dataset(st: SessionState, window_sec: float, hop_sec: float, fs_target: float):
352
+ class_names = sorted([k for k in st.data.keys() if k.strip()])
353
+ if len(class_names) < 2:
354
+ return [], np.zeros((0,), np.int64), class_names
355
+
356
+ seqs = []
357
+ ys = []
358
+ for lab_idx, lab in enumerate(class_names):
359
+ for item in st.data[lab]:
360
+ ts = item["t"].astype(np.float32)
361
+ A = item["a"].astype(np.float32)
362
+
363
+ ts2, A2 = resample_linear(ts, A, fs_target)
364
+
365
+ hp = GravityHighPass(alpha=CFG.hp_alpha)
366
+ Ad = np.stack([hp.step(A2[i]) for i in range(len(A2))], axis=0)
367
+
368
+ an = np.linalg.norm(Ad, axis=1, keepdims=True)
369
+ X = np.concatenate([Ad, an], axis=1) # (T, 4)
370
+
371
+ win = int(round(window_sec * fs_target))
372
+ hop = int(round(hop_sec * fs_target))
373
+ if len(X) < win:
374
+ continue
375
+
376
+ for s in range(0, len(X) - win + 1, hop):
377
+ seqs.append(X[s:s + win])
378
+ ys.append(lab_idx)
379
+
380
+ return seqs, np.array(ys, dtype=np.int64), class_names
381
+
382
+
383
+ def train_click(window_sec, hop_sec, fs_target, feat_mode, request: gr.Request):
384
+ sid = _get_sid_from_request(request)
385
+ st = get_state(sid)
386
+
387
+ window_sec = float(window_sec); hop_sec = float(hop_sec); fs_target = float(fs_target)
388
+
389
+ with st.lock:
390
+ seqs, y, class_names = build_dataset(st, window_sec, hop_sec, fs_target)
391
+
392
+ if len(seqs) < 12 or len(class_names) < 2:
393
+ return ui_status_for(st), {}, "データ不足(2ラベル以上、各数回〜推奨)", ""
394
+
395
+ idx = np.arange(len(seqs))
396
+ try:
397
+ tr_idx, va_idx = train_test_split(idx, test_size=0.25, random_state=0, stratify=y)
398
+ except Exception:
399
+ tr_idx, va_idx = train_test_split(idx, test_size=0.25, random_state=0)
400
+
401
+ Xtr_all = np.concatenate([seqs[i] for i in tr_idx], axis=0)
402
+ mu = Xtr_all.mean(axis=0, keepdims=True).astype(np.float32)
403
+ sd = (Xtr_all.std(axis=0, keepdims=True) + 1e-6).astype(np.float32)
404
+
405
+ def norm_seq(seg):
406
+ return (seg - mu) / sd
407
+
408
+ cand_res = [80, 120]
409
+ cand_sr = [0.8, 1.0]
410
+ cand_leak = [0.2, 0.5, 0.8]
411
+ cand_ridge = [1e-3]
412
+
413
+ best_acc = -1.0
414
+ best_pack = None
415
+ logs = []
416
+
417
+ for res_size in cand_res:
418
+ for sr in cand_sr:
419
+ for leak in cand_leak:
420
+ for ridge in cand_ridge:
421
+ esn = ESNClassifier(
422
+ in_dim=4,
423
+ res_size=int(res_size),
424
+ spectral_radius=float(sr),
425
+ leak=float(leak),
426
+ ridge=float(ridge),
427
+ seed=0
428
+ )
429
+
430
+ Xtr = []
431
+ for i in tr_idx:
432
+ feat = make_window_feature(esn, norm_seq(seqs[i]), mode=feat_mode)
433
+ Xtr.append(feat)
434
+ Xtr = np.stack(Xtr, axis=0).astype(np.float32)
435
+
436
+ esn.fit(Xtr, y[tr_idx], class_names)
437
+
438
+ correct = 0
439
+ for i in va_idx:
440
+ feat = make_window_feature(esn, norm_seq(seqs[i]), mode=feat_mode)
441
+ p = esn.predict_proba(feat)
442
+ pred = int(np.argmax(p))
443
+ correct += (pred == int(y[i]))
444
+ acc = correct / max(1, len(va_idx))
445
+
446
+ logs.append(f"res={res_size}, sr={sr}, leak={leak}, ridge={ridge} -> val_acc={acc:.3f}")
447
+
448
+ if acc > best_acc:
449
+ best_acc = acc
450
+ best_pack = (int(res_size), float(sr), float(leak), float(ridge), esn)
451
+
452
+ res_size, sr, leak, ridge, esn = best_pack
453
+
454
+ with st.lock:
455
+ st.trained = True
456
+ st.pp_mean, st.pp_std = mu, sd
457
+ st.esn_model = esn
458
+ st.train_cfg = {
459
+ "window_sec": window_sec,
460
+ "hop_sec": hop_sec,
461
+ "fs_target": fs_target,
462
+ "mode": feat_mode,
463
+ "res_size": res_size,
464
+ "spectral_radius": sr,
465
+ "leak": leak,
466
+ "ridge": ridge,
467
+ "val_acc": float(best_acc),
468
+ "classes": class_names,
469
+ }
470
+
471
+ tail = "\n".join(logs[-12:])
472
+ return ui_status_for(st), st.train_cfg, f"学習完了: val_acc={best_acc:.3f}", tail
473
+
474
+
475
+ # =========================
476
+ # Inference
477
+ # =========================
478
+ def infer_step(st: SessionState):
479
+ if (not st.trained) or (st.esn_model is None) or (st.pp_mean is None) or (st.pp_std is None):
480
+ return "(not trained)", 0.0, {}
481
+
482
+ fs = float(st.train_cfg["fs_target"])
483
+ win = int(round(float(st.train_cfg["window_sec"]) * fs))
484
+ if len(st.stream_t) < win + 2:
485
+ return "(buffering)", 0.0, {}
486
+
487
+ ts = np.array(st.stream_t, dtype=np.float32)
488
+ A = np.array(st.stream_a, dtype=np.float32)
489
+
490
+ t_end = ts[-1]
491
+ t_start = max(ts[0], t_end - (float(st.train_cfg["window_sec"]) + 0.4))
492
+ m = ts >= t_start
493
+ ts2, A2 = resample_linear(ts[m], A[m], fs)
494
+ if len(ts2) < win:
495
+ return "(buffering)", 0.0, {}
496
+ A2 = A2[-win:]
497
+
498
+ hp = GravityHighPass(alpha=CFG.hp_alpha)
499
+ Ad = np.stack([hp.step(A2[i]) for i in range(len(A2))], axis=0)
500
+ an = np.linalg.norm(Ad, axis=1, keepdims=True)
501
+ X = np.concatenate([Ad, an], axis=1).astype(np.float32)
502
+
503
+ Xn = (X - st.pp_mean) / st.pp_std
504
+ feat = make_window_feature(st.esn_model, Xn, mode=st.train_cfg["mode"])
505
+ p = st.esn_model.predict_proba(feat)
506
+
507
+ i = int(np.argmax(p))
508
+ conf = float(p[i])
509
+ lab = st.train_cfg["classes"][i]
510
+
511
+ prev_lab = st.infer_last_label
512
+ st.infer_last_label = lab
513
+ st.infer_last_conf = conf
514
+ if (lab != prev_lab) and (lab not in ["(buffering)", "(not trained)"]):
515
+ st.infer_pred_log.append((float(t_end), lab, conf))
516
+ st.infer_pred_log = st.infer_pred_log[-500:]
517
+
518
+ info = {"probs": {st.train_cfg["classes"][j]: float(p[j]) for j in range(len(p))}}
519
+ return lab, conf, info
520
+
521
+
522
+ def infer_start(request: gr.Request):
523
+ sid = _get_sid_from_request(request)
524
+ st = get_state(sid)
525
+ with st.lock:
526
+ if not st.trained:
527
+ return "学習してから推論してください", "<div style='font-size:24px;font-weight:800;opacity:.6'>-</div>", {}, ui_status_for(st)
528
+ st.infer_running = True
529
+ st.infer_pred_log = []
530
+ st.infer_last_label = ""
531
+ st.infer_last_conf = 0.0
532
+ return "推論: ON", "<div style='font-size:24px;font-weight:800;opacity:.6'>-</div>", {}, ui_status_for(st)
533
+
534
+
535
+ def infer_stop(request: gr.Request):
536
+ sid = _get_sid_from_request(request)
537
+ st = get_state(sid)
538
+ with st.lock:
539
+ st.infer_running = False
540
+ return "推論: OFF", "<div style='font-size:24px;font-weight:800;opacity:.6'>-</div>", {}, ui_status_for(st)
541
+
542
+
543
+ def infer_tick(request: gr.Request):
544
+ sid = _get_sid_from_request(request)
545
+ st = get_state(sid)
546
+
547
+ with st.lock:
548
+ if not st.infer_running:
549
+ return gr.update(), gr.update(), gr.update()
550
+ lab, conf, info = infer_step(st)
551
+ stline = ui_status_for(st)
552
+
553
+ pred_html = (
554
+ f"<div style='padding:10px 12px;border:1px solid #ddd;border-radius:12px;background:#fff'>"
555
+ f"<div style='font-size:30px;font-weight:900;line-height:1.1'>{lab}</div>"
556
+ f"<div style='font-size:13px;opacity:.85'>conf={conf:.2f}</div>"
557
+ f"</div>"
558
+ )
559
+ probs = info.get("probs", {})
560
+ return pred_html, probs, stline
561
+
562
+
563
+ def chat_tick(request: gr.Request):
564
+ sid = _get_sid_from_request(request)
565
+ st = get_state(sid)
566
+ with st.lock:
567
+ if not st.infer_running:
568
+ big = "<div style='font-size:22px;font-weight:800;opacity:.6'>推論がOFFです</div>"
569
+ log_md = format_pred_log_md(st)
570
+ return big, log_md
571
+ big = (
572
+ f"<div style='padding:12px 14px;border:1px solid #ddd;border-radius:14px;background:#fff'>"
573
+ f"<div style='font-size:34px;font-weight:900;line-height:1.05'>{st.infer_last_label or '-'}</div>"
574
+ f"<div style='font-size:14px;opacity:.85'>conf={st.infer_last_conf:.2f}</div>"
575
+ f"</div>"
576
+ )
577
+ log_md = format_pred_log_md(st)
578
+ return big, log_md
579
+
580
+
581
+ # =========================
582
+ # JS UI + Boot
583
+ # - sid is stored in localStorage and sent in JSON payload
584
+ # =========================
585
+
586
+ # ── UI変更: SENSOR_UI をフェミニンデザインに全面リデザイン ──
587
+ # パステル背景・中央寄せ・やわらかい角丸ボタン・透明感・余白多め
588
+ SENSOR_UI = r"""
589
+ <div class="sensor-hero">
590
+ <div class="sensor-hero__icon">&#9752;</div>
591
+ <div class="sensor-hero__title">Motion Sensor</div>
592
+ <div class="sensor-hero__subtitle">スマホを振って、動きを学習させよう</div>
593
+ <div class="sensor-hero__buttons">
594
+ <button id="btn_perm" class="hero-btn hero-btn--outline" type="button">PERMISSION</button>
595
+ <button id="btn_start" class="hero-btn hero-btn--dark" type="button">START</button>
596
+ <button id="btn_stop" class="hero-btn hero-btn--outline" type="button">STOP</button>
597
+ <button id="btn_reset" class="hero-btn hero-btn--dark" type="button">RESET</button>
598
+ </div>
599
+ <div class="sensor-hero__status">
600
+ <span id="sensor_status">status: idle</span>
601
+ </div>
602
+ <div class="sensor-hero__share">
603
+ URL: <span id="share_url" class="mono"></span>
604
+ </div>
605
+ </div>
606
+ """
607
+
608
+ JS_BOOT = r"""
609
+ () => {
610
+ const setStatus = (s) => {
611
+ const el = document.getElementById('sensor_status');
612
+ if (el) el.textContent = 'status: ' + s;
613
+ };
614
+
615
+ const setShareUrl = () => {
616
+ const el = document.getElementById('share_url');
617
+ if (!el) return;
618
+ el.textContent = window.location.href;
619
+ };
620
+
621
+ const genSid = () => {
622
+ if (crypto && crypto.randomUUID) return crypto.randomUUID();
623
+ const r = () => Math.floor(Math.random() * 1e9).toString(16);
624
+ return `${Date.now().toString(16)}-${r()}-${r()}-${r()}`;
625
+ };
626
+
627
+ const getSid = () => {
628
+ try{
629
+ const k = "sid_v1";
630
+ let sid = localStorage.getItem(k);
631
+ if(!sid){ sid = genSid(); localStorage.setItem(k, sid); }
632
+ return sid;
633
+ }catch(e){
634
+ return genSid();
635
+ }
636
+ };
637
+
638
+ // Gradio root (subpath/iframeでも壊れにくい)
639
+ const apiUrl = (path) => {
640
+ const root = (window.gradio_config && window.gradio_config.root) ? window.gradio_config.root : '';
641
+ return `${root}${path}`;
642
+ };
643
+
644
+ const sid = getSid();
645
+
646
+ let running=false, buf=[], t0=null, timer=null;
647
+ let accel=null, dmHandler=null;
648
+
649
+ const pushSample = (ax,ay,az) => {
650
+ if(t0===null) t0=performance.now();
651
+ const t=(performance.now()-t0)/1000.0;
652
+ buf.push({t, ax:ax||0, ay:ay||0, az:az||0});
653
+ if(buf.length>600) buf=buf.slice(-600);
654
+ };
655
+
656
+ const postJson = async (path, payloadObj) => {
657
+ try{
658
+ payloadObj = payloadObj || {};
659
+ payloadObj.sid = sid; // <-- include sid in body
660
+ const res = await fetch(apiUrl(path), {
661
+ method: 'POST',
662
+ headers: {'Content-Type':'application/json'},
663
+ body: JSON.stringify(payloadObj),
664
+ credentials: 'include'
665
+ });
666
+ if(!res.ok){
667
+ setStatus(`ERR: ${path} http ${res.status}`);
668
+ return false;
669
+ }
670
+ return true;
671
+ }catch(e){
672
+ setStatus(`ERR: fetch(${path}) failed`);
673
+ return false;
674
+ }
675
+ };
676
+
677
+ const flush = async () => {
678
+ if(!running || buf.length===0) return;
679
+ const samples = buf;
680
+ buf = [];
681
+ await postJson('/api/ingest', {samples});
682
+ };
683
+
684
+ const startDeviceMotion = () => {
685
+ dmHandler = (e) => {
686
+ if(!running) return;
687
+ const acc = e.accelerationIncludingGravity || e.acceleration;
688
+ if(!acc) return;
689
+ pushSample(acc.x, acc.y, acc.z);
690
+ };
691
+ window.addEventListener('devicemotion', dmHandler, {passive:true});
692
+ };
693
+
694
+ const startGeneric = () => {
695
+ if(!('Accelerometer' in window)) return false;
696
+ try{
697
+ accel = new Accelerometer({frequency: 50});
698
+ accel.addEventListener('reading', ()=>{ if(running) pushSample(accel.x, accel.y, accel.z); }, {passive:true});
699
+ accel.addEventListener('error', ()=>{ try{accel.stop();}catch(e){} accel=null; startDeviceMotion(); });
700
+ accel.start();
701
+ return true;
702
+ }catch(e){
703
+ accel=null;
704
+ return false;
705
+ }
706
+ };
707
+
708
+ const requestPerm = async () => {
709
+ try{
710
+ if(typeof DeviceMotionEvent!=='undefined' && typeof DeviceMotionEvent.requestPermission==='function'){
711
+ const res = await DeviceMotionEvent.requestPermission();
712
+ setStatus('permission: '+res + ' / sid=' + sid.slice(0,8));
713
+ } else {
714
+ setStatus('permission: not-needed / sid=' + sid.slice(0,8));
715
+ }
716
+ }catch(e){
717
+ setStatus('permission error');
718
+ }
719
+ };
720
+
721
+ const start = () => {
722
+ if(running) return;
723
+ running=true; buf=[]; t0=null;
724
+ if(!startGeneric()) startDeviceMotion();
725
+ timer=setInterval(flush, 200);
726
+ setStatus('running / sid=' + sid.slice(0,8));
727
+ };
728
+
729
+ const stop = () => {
730
+ if(!running) return;
731
+ running=false;
732
+ if(accel){ try{accel.stop();}catch(e){} accel=null; }
733
+ if(dmHandler){ window.removeEventListener('devicemotion', dmHandler); dmHandler=null; }
734
+ if(timer){ clearInterval(timer); timer=null; }
735
+ flush();
736
+ setStatus('stopped / sid=' + sid.slice(0,8));
737
+ };
738
+
739
+ const reset = async () => {
740
+ stop();
741
+ await postJson('/api/reset', {});
742
+ setStatus('reset done / sid=' + sid.slice(0,8));
743
+ };
744
+
745
+ const bind = () => {
746
+ const p=document.getElementById('btn_perm');
747
+ const s=document.getElementById('btn_start');
748
+ const x=document.getElementById('btn_stop');
749
+ const r=document.getElementById('btn_reset');
750
+ if(!p || !s || !x || !r){ setTimeout(bind, 300); return; }
751
+ p.onclick=requestPerm;
752
+ s.onclick=start;
753
+ x.onclick=stop;
754
+ r.onclick=reset;
755
+ setStatus('ready / sid=' + sid.slice(0,8));
756
+ setShareUrl();
757
+ };
758
+ bind();
759
+ }
760
+ """
761
+
762
+ # ── UI変更: CSS をフェミニン・パステルデザインに全面リデザイン ──
763
+ # くすみピンク / ミント / ラベンダー / 低彩度 / 黒不使用 / 透明感 / 余白多め
764
+ CSS = """
765
+ /* ========================================
766
+ グローバル: フェミニン・パステルテーマ
767
+ ======================================== */
768
+ html {
769
+ scroll-behavior: smooth !important;
770
+ -webkit-overflow-scrolling: touch !important;
771
+ }
772
+
773
+ /* Gradio コンテナ: 淡いグラデーション背景 */
774
+ .gradio-container {
775
+ background: linear-gradient(175deg, #fdf2f8 0%, #faf5ff 35%, #f0fdf4 70%, #fdf2f8 100%) !important;
776
+ color: #1a1a1a !important;
777
+ font-family: 'Inter', 'Hiragino Kaku Gothic ProN', 'Noto Sans JP', -apple-system, BlinkMacSystemFont, sans-serif !important;
778
+ font-weight: 400 !important;
779
+ max-width: 100% !important;
780
+ padding: 0 !important;
781
+ min-height: 100vh !important;
782
+ }
783
+
784
+ /* フッター非表示 */
785
+ footer { display: none !important; }
786
+
787
+ /* ========================================
788
+ センサーヒーローセクション
789
+ ======================================== */
790
+ .sensor-hero {
791
+ min-height: 65vh;
792
+ display: flex;
793
+ flex-direction: column;
794
+ align-items: center;
795
+ justify-content: center;
796
+ text-align: center;
797
+ padding: 56px 24px 48px 24px;
798
+ background: linear-gradient(170deg,
799
+ rgba(253,242,248,0.9) 0%,
800
+ rgba(250,245,255,0.85) 40%,
801
+ rgba(240,253,244,0.8) 100%);
802
+ margin-bottom: 8px;
803
+ }
804
+
805
+ .sensor-hero__icon {
806
+ font-size: 36px;
807
+ margin-bottom: 16px;
808
+ opacity: 0.6;
809
+ filter: grayscale(30%);
810
+ }
811
+
812
+ .sensor-hero__title {
813
+ font-family: 'Cormorant Garamond', 'Georgia', 'Times New Roman', serif;
814
+ font-size: clamp(30px, 8vw, 48px);
815
+ font-weight: 400;
816
+ letter-spacing: 0.08em;
817
+ color: #1a1a1a;
818
+ margin-bottom: 10px;
819
+ line-height: 1.15;
820
+ text-transform: uppercase;
821
+ }
822
+
823
+ .sensor-hero__subtitle {
824
+ font-size: clamp(13px, 3.2vw, 16px);
825
+ font-weight: 400;
826
+ color: #333333;
827
+ margin-bottom: 44px;
828
+ letter-spacing: 0.03em;
829
+ line-height: 1.6;
830
+ }
831
+
832
+ .sensor-hero__buttons {
833
+ display: flex;
834
+ flex-wrap: wrap;
835
+ gap: 10px;
836
+ justify-content: center;
837
+ margin-bottom: 36px;
838
+ max-width: 360px;
839
+ }
840
+
841
+ .sensor-hero__status {
842
+ font-size: 12px;
843
+ color: #333333;
844
+ font-family: 'SF Mono', 'Fira Code', ui-monospace, monospace;
845
+ margin-bottom: 6px;
846
+ letter-spacing: 0.02em;
847
+ }
848
+
849
+ .sensor-hero__share {
850
+ font-size: 11px;
851
+ color: #333333;
852
+ word-break: break-all;
853
+ max-width: 85vw;
854
+ }
855
+ .sensor-hero__share .mono {
856
+ font-family: 'SF Mono', 'Fira Code', ui-monospace, monospace;
857
+ }
858
+
859
+ /* ========================================
860
+ ヒーローボタン: VIEW MORE風 / セリフ体 / シャープ
861
+ ======================================== */
862
+ .hero-btn {
863
+ flex: 1 1 calc(50% - 5px);
864
+ box-sizing: border-box;
865
+ padding: 15px 10px;
866
+ border-radius: 0;
867
+ font-family: 'Cormorant Garamond', 'Georgia', 'Times New Roman', 'YuMincho', serif;
868
+ font-size: 13px;
869
+ font-weight: 500;
870
+ letter-spacing: 0.18em;
871
+ text-transform: uppercase;
872
+ text-align: center;
873
+ cursor: pointer;
874
+ transition: all 0.3s ease;
875
+ touch-action: manipulation;
876
+ -webkit-tap-highlight-color: transparent;
877
+ }
878
+
879
+ /* 白背景 + 細線ボーダー(左のボタン) */
880
+ .hero-btn--outline {
881
+ background: #ffffff;
882
+ color: #555555;
883
+ border: 1px solid #aaaaaa;
884
+ }
885
+ .hero-btn--outline:active {
886
+ background: #f5f5f5;
887
+ border-color: #888888;
888
+ }
889
+
890
+ /* ダーク背景(右のボタン) */
891
+ .hero-btn--dark {
892
+ background: #3a3a3a;
893
+ color: #d8d8d8;
894
+ border: 1px solid #3a3a3a;
895
+ }
896
+ .hero-btn--dark:active {
897
+ background: #4a4a4a;
898
+ }
899
+
900
+ /* ========================================
901
+ タブナビゲーション: ピル型パステル
902
+ ======================================== */
903
+ div.tab-nav {
904
+ background: rgba(255,255,255,0.8) !important;
905
+ border: 1px solid rgba(0,0,0,0.06) !important;
906
+ border-radius: 22px !important;
907
+ padding: 5px !important;
908
+ margin: 16px 16px 20px 16px !important;
909
+ display: flex !important;
910
+ justify-content: center !important;
911
+ gap: 3px !important;
912
+ box-shadow: 0 2px 12px rgba(107,91,123,0.06) !important;
913
+ backdrop-filter: blur(8px) !important;
914
+ -webkit-backdrop-filter: blur(8px) !important;
915
+ }
916
+ div.tab-nav button {
917
+ background: transparent !important;
918
+ color: #444444 !important;
919
+ border: none !important;
920
+ border-radius: 18px !important;
921
+ padding: 10px 18px !important;
922
+ font-family: 'Cormorant Garamond', 'Georgia', 'Times New Roman', serif !important;
923
+ font-size: 14px !important;
924
+ font-weight: 500 !important;
925
+ transition: all 0.25s ease !important;
926
+ letter-spacing: 0.1em !important;
927
+ }
928
+ div.tab-nav button.selected {
929
+ background: rgba(244,196,212,0.35) !important;
930
+ color: #222222 !important;
931
+ box-shadow: 0 1px 8px rgba(244,196,212,0.2) !important;
932
+ }
933
+
934
+ /* ========================================
935
+ タブコンテンツ: 二重フレーム(黒+グレーずらし)
936
+ ======================================== */
937
+ .tabitem {
938
+ background: transparent !important;
939
+ border: none !important;
940
+ }
941
+ .tabitem > div {
942
+ position: relative !important;
943
+ background: #ffffff !important;
944
+ border-radius: 0 !important;
945
+ padding: 36px 24px !important;
946
+ margin: 20px 22px 32px 22px !important;
947
+ border: 1.25px solid #888888 !important;
948
+ box-shadow: 10px 10px 0px 0px #c8c8c8 !important;
949
+ backdrop-filter: none !important;
950
+ -webkit-backdrop-filter: none !important;
951
+ }
952
+
953
+ /* ========================================
954
+ Gradioコンポーネントのスタイリング
955
+ ======================================== */
956
+
957
+ /* ラベル */
958
+ label span, .label-wrap span {
959
+ color: #222222 !important;
960
+ font-family: 'Cormorant Garamond', 'Georgia', 'Times New Roman', serif !important;
961
+ font-weight: 500 !important;
962
+ font-size: 14px !important;
963
+ letter-spacing: 0.06em !important;
964
+ }
965
+
966
+ /* テキスト入力 / Dropdown */
967
+ input[type="text"], textarea, select {
968
+ background: rgba(255,255,255,0.7) !important;
969
+ border: 1px solid rgba(200,191,224,0.3) !important;
970
+ border-radius: 16px !important;
971
+ color: #1a1a1a !important;
972
+ font-weight: 400 !important;
973
+ }
974
+ input[type="text"]:focus, textarea:focus {
975
+ border-color: rgba(244,196,212,0.5) !important;
976
+ box-shadow: 0 0 0 3px rgba(244,196,212,0.15) !important;
977
+ outline: none !important;
978
+ }
979
+
980
+ /* Slider number input: 大人っぽく角ばった四角 */
981
+ input[type="number"] {
982
+ background: #ffffff !important;
983
+ border: 1.25px solid #888888 !important;
984
+ border-radius: 0 !important;
985
+ color: #1a1a1a !important;
986
+ font-family: 'Cormorant Garamond', 'Georgia', serif !important;
987
+ font-weight: 500 !important;
988
+ font-size: 13px !important;
989
+ letter-spacing: 0.05em !important;
990
+ text-align: center !important;
991
+ padding: 4px 6px !important;
992
+ box-shadow: 3px 3px 0px 0px #c8c8c8 !important;
993
+ outline: none !important;
994
+ -moz-appearance: textfield !important;
995
+ }
996
+ input[type="number"]:focus {
997
+ border-color: #555555 !important;
998
+ box-shadow: 4px 4px 0px 0px #aaaaaa !important;
999
+ outline: none !important;
1000
+ }
1001
+ input[type="number"]::-webkit-inner-spin-button,
1002
+ input[type="number"]::-webkit-outer-spin-button {
1003
+ -webkit-appearance: none !important;
1004
+ margin: 0 !important;
1005
+ }
1006
+
1007
+ /* Dropdown: 全体リセット */
1008
+ [data-testid="dropdown"] {
1009
+ background: transparent !important;
1010
+ border: none !important;
1011
+ box-shadow: none !important;
1012
+ border-radius: 0 !important;
1013
+ }
1014
+ [data-testid="dropdown"] > div,
1015
+ [data-testid="dropdown"] .wrap,
1016
+ [data-testid="dropdown"] .wrap-inner,
1017
+ [data-testid="dropdown"] .secondary-wrap,
1018
+ [data-testid="dropdown"] input,
1019
+ [data-testid="dropdown"] .multiselect {
1020
+ background: #ffffff !important;
1021
+ background-color: #ffffff !important;
1022
+ border: none !important;
1023
+ border-radius: 0 !important;
1024
+ box-shadow: none !important;
1025
+ outline: none !important;
1026
+ }
1027
+ /* 入力ラッパーのみ細い黒線で囲む */
1028
+ [data-testid="dropdown"] .wrap,
1029
+ [data-testid="dropdown"] .secondary-wrap {
1030
+ border: 1px solid #1a1a1a !important;
1031
+ box-shadow: 3px 3px 0px 0px #cccccc !important;
1032
+ padding: 8px 10px !important;
1033
+ }
1034
+ /* 子要素の文字色 */
1035
+ [data-testid="dropdown"] *:not(ul):not(ul *) {
1036
+ background: #ffffff !important;
1037
+ background-color: #ffffff !important;
1038
+ color: #1a1a1a !important;
1039
+ border-radius: 0 !important;
1040
+ }
1041
+
1042
+ /* Dropdown選択肢リスト: 黒背景・白文字 */
1043
+ ul.options,
1044
+ ul.options li,
1045
+ .options,
1046
+ .options .item,
1047
+ .secondary-wrap .item,
1048
+ .secondary-wrap ul li {
1049
+ background: #1a1a1a !important;
1050
+ color: #ffffff !important;
1051
+ border-radius: 0 !important;
1052
+ }
1053
+ ul.options li:hover,
1054
+ .options .item:hover,
1055
+ .secondary-wrap .item:hover {
1056
+ background: #333333 !important;
1057
+ color: #ffffff !important;
1058
+ }
1059
+ ul.options li.selected,
1060
+ .options .item.active {
1061
+ background: #555555 !important;
1062
+ color: #ffffff !important;
1063
+ }
1064
+
1065
+ /* Slider: ラグジュアリー仕様 */
1066
+ input[type="range"] {
1067
+ -webkit-appearance: none !important;
1068
+ appearance: none !important;
1069
+ height: 3px !important;
1070
+ background: linear-gradient(90deg,
1071
+ #e8c8d4 0%,
1072
+ #d4b8e0 40%,
1073
+ #b8d4e8 100%) !important;
1074
+ border-radius: 0 !important;
1075
+ outline: none !important;
1076
+ cursor: pointer !important;
1077
+ overflow: visible !important;
1078
+ margin: 12px 0 !important;
1079
+ }
1080
+ input[type="range"]::-webkit-slider-thumb {
1081
+ -webkit-appearance: none !important;
1082
+ appearance: none !important;
1083
+ width: 14px !important;
1084
+ height: 14px !important;
1085
+ background: #3a3a3a !important;
1086
+ border: 1.5px solid #888888 !important;
1087
+ border-radius: 0 !important;
1088
+ transform: rotate(45deg) !important;
1089
+ cursor: pointer !important;
1090
+ box-shadow: 2px 2px 4px rgba(0,0,0,0.2) !important;
1091
+ margin-top: -6px !important;
1092
+ position: relative !important;
1093
+ }
1094
+ input[type="range"]::-moz-range-thumb {
1095
+ width: 14px !important;
1096
+ height: 14px !important;
1097
+ background: #3a3a3a !important;
1098
+ border: 1.5px solid #888888 !important;
1099
+ border-radius: 0 !important;
1100
+ transform: rotate(45deg) !important;
1101
+ cursor: pointer !important;
1102
+ }
1103
+ input[type="range"]::-webkit-slider-runnable-track {
1104
+ height: 3px !important;
1105
+ background: linear-gradient(90deg,
1106
+ #e8c8d4 0%,
1107
+ #d4b8e0 40%,
1108
+ #b8d4e8 100%) !important;
1109
+ border-radius: 0 !important;
1110
+ overflow: visible !important;
1111
+ }
1112
+
1113
+ /* Sliderブロック全体: 角ばったコンテナ */
1114
+ [data-testid="slider"] {
1115
+ background: #fafafa !important;
1116
+ border: 1.25px solid #cccccc !important;
1117
+ border-radius: 0 !important;
1118
+ padding: 12px 14px 16px 14px !important;
1119
+ box-shadow: 3px 3px 0px 0px #d8d8d8 !important;
1120
+ overflow: visible !important;
1121
+ }
1122
+ [data-testid="slider"] > div,
1123
+ [data-testid="slider"] .wrap,
1124
+ [data-testid="slider"] .wrap-inner {
1125
+ overflow: visible !important;
1126
+ }
1127
+ [data-testid="slider"] .label-wrap span,
1128
+ [data-testid="slider"] label span {
1129
+ font-family: 'Cormorant Garamond', 'Georgia', serif !important;
1130
+ font-size: 11px !important;
1131
+ letter-spacing: 0.12em !important;
1132
+ text-transform: uppercase !important;
1133
+ color: #555555 !important;
1134
+ }
1135
+ /* sliderのrefreshボタン(リセットアイコン)を角ばりに */
1136
+ [data-testid="slider"] button {
1137
+ border-radius: 0 !important;
1138
+ border: 1.25px solid #aaaaaa !important;
1139
+ background: #f0f0f0 !important;
1140
+ padding: 4px 6px !important;
1141
+ }
1142
+
1143
+ /* Radio */
1144
+ .gr-radio-row label, [data-testid="radio-group"] label {
1145
+ color: #222222 !important;
1146
+ font-weight: 400 !important;
1147
+ }
1148
+
1149
+ /* JSON表示 */
1150
+ .json-holder, [data-testid="json"] {
1151
+ background: rgba(255,255,255,0.5) !important;
1152
+ border-radius: 18px !important;
1153
+ border: 1px solid rgba(200,191,224,0.15) !important;
1154
+ }
1155
+
1156
+ /* Textbox */
1157
+ textarea {
1158
+ background: rgba(255,255,255,0.6) !important;
1159
+ color: #1a1a1a !important;
1160
+ border-radius: 16px !important;
1161
+ font-family: 'SF Mono', 'Fira Code', ui-monospace, monospace !important;
1162
+ font-size: 12px !important;
1163
+ }
1164
+
1165
+ /* Markdown */
1166
+ .prose, .markdown-text, .md {
1167
+ color: #1a1a1a !important;
1168
+ }
1169
+ .prose h2, .prose h3 {
1170
+ font-family: 'Cormorant Garamond', 'Georgia', 'Times New Roman', serif !important;
1171
+ color: #222222 !important;
1172
+ font-weight: 400 !important;
1173
+ letter-spacing: 0.08em !important;
1174
+ }
1175
+ .prose table {
1176
+ color: #1a1a1a !important;
1177
+ }
1178
+ .prose table th {
1179
+ color: #222222 !important;
1180
+ font-weight: 500 !important;
1181
+ background: rgba(244,196,212,0.1) !important;
1182
+ }
1183
+ .prose table td {
1184
+ border-color: rgba(200,191,224,0.2) !important;
1185
+ }
1186
+
1187
+ /* ========================================
1188
+ Gradioボタン: VIEW MORE風 / セリフ体 / シャープ
1189
+ ======================================== */
1190
+ button[class*="primary"], button[class*="secondary"],
1191
+ button.lg {
1192
+ border-radius: 0 !important;
1193
+ font-family: 'Cormorant Garamond', 'Georgia', 'Times New Roman', 'YuMincho', serif !important;
1194
+ font-weight: 500 !important;
1195
+ font-size: 13px !important;
1196
+ letter-spacing: 0.18em !important;
1197
+ text-transform: uppercase !important;
1198
+ padding: 16px 28px !important;
1199
+ transition: all 0.3s ease !important;
1200
+ touch-action: manipulation !important;
1201
+ -webkit-tap-highlight-color: transparent !important;
1202
+ }
1203
+
1204
+ /* Primary ボタン: ダーク背景 */
1205
+ button[class*="primary"] {
1206
+ background: #3a3a3a !important;
1207
+ color: #d8d8d8 !important;
1208
+ border: 1px solid #3a3a3a !important;
1209
+ box-shadow: none !important;
1210
+ }
1211
+ button[class*="primary"]:hover {
1212
+ background: #4a4a4a !important;
1213
+ transform: none !important;
1214
+ box-shadow: none !important;
1215
+ }
1216
+ button[class*="primary"]:active {
1217
+ background: #555555 !important;
1218
+ }
1219
+
1220
+ /* Secondary ボタン: 白背景 + 細線ボーダー */
1221
+ button[class*="secondary"] {
1222
+ background: #ffffff !important;
1223
+ color: #555555 !important;
1224
+ border: 1px solid #aaaaaa !important;
1225
+ box-shadow: none !important;
1226
+ }
1227
+ button[class*="secondary"]:hover {
1228
+ background: #f5f5f5 !important;
1229
+ border-color: #888888 !important;
1230
+ transform: none !important;
1231
+ box-shadow: none !important;
1232
+ }
1233
+ button[class*="secondary"]:active {
1234
+ background: #eeeeee !important;
1235
+ }
1236
+
1237
+ /* ========================================
1238
+ セクションタイトル
1239
+ ======================================== */
1240
+ .section-title {
1241
+ text-align: center !important;
1242
+ padding: 20px 16px 4px 16px !important;
1243
+ }
1244
+ .section-title h2 {
1245
+ font-family: 'Cormorant Garamond', 'Georgia', 'Times New Roman', serif !important;
1246
+ font-size: clamp(16px, 4vw, 22px) !important;
1247
+ font-weight: 400 !important;
1248
+ color: #222222 !important;
1249
+ letter-spacing: 0.15em !important;
1250
+ text-transform: uppercase !important;
1251
+ }
1252
+
1253
+ /* ========================================
1254
+ WORKFLOW以下: 背景を統一して視認性UP
1255
+ ======================================== */
1256
+ .section-title,
1257
+ .section-title ~ * {
1258
+ background-color: #ffffff !important;
1259
+ }
1260
+
1261
+ /* ========================================
1262
+ レ��ポンシブ: PC = 中央固定幅
1263
+ ======================================== */
1264
+ @media (min-width: 768px) {
1265
+ .gradio-container > .main,
1266
+ .gradio-container > div > .main {
1267
+ max-width: 480px !important;
1268
+ margin: 0 auto !important;
1269
+ }
1270
+ .sensor-hero {
1271
+ min-height: 55vh;
1272
+ }
1273
+ .tabitem > div {
1274
+ margin: 20px auto 32px auto !important;
1275
+ max-width: 440px !important;
1276
+ }
1277
+ div.tab-nav {
1278
+ max-width: 440px !important;
1279
+ margin: 16px auto 20px auto !important;
1280
+ }
1281
+ }
1282
+
1283
+ /* ========================================
1284
+ スマホ特化: タッチ最適化
1285
+ ======================================== */
1286
+ @media (max-width: 767px) {
1287
+ .sensor-hero {
1288
+ min-height: 70vh;
1289
+ padding: 48px 20px 40px 20px;
1290
+ }
1291
+ .sensor-hero__buttons {
1292
+ width: 100%;
1293
+ max-width: 300px;
1294
+ }
1295
+ .hero-btn {
1296
+ flex: 1 1 calc(50% - 5px);
1297
+ min-width: 130px;
1298
+ padding: 14px 10px;
1299
+ font-size: 13px;
1300
+ }
1301
+ .section-title {
1302
+ padding: 24px 16px 4px 16px !important;
1303
+ }
1304
+ div.tab-nav {
1305
+ margin: 12px 12px 16px 12px !important;
1306
+ padding: 4px !important;
1307
+ }
1308
+ div.tab-nav button {
1309
+ padding: 9px 12px !important;
1310
+ font-size: 13px !important;
1311
+ }
1312
+ .tabitem > div {
1313
+ margin: 16px 14px 28px 14px !important;
1314
+ padding: 28px 16px !important;
1315
+ box-shadow: 8px 8px 0px 0px #c8c8c8 !important;
1316
+ }
1317
+ button[class*="primary"], button[class*="secondary"],
1318
+ button.lg {
1319
+ width: 100% !important;
1320
+ padding: 15px 20px !important;
1321
+ font-size: 15px !important;
1322
+ }
1323
+ label span, .label-wrap span {
1324
+ font-size: 13px !important;
1325
+ }
1326
+ input[type="text"], input[type="number"], textarea, select {
1327
+ font-size: 16px !important;
1328
+ }
1329
+ /* Rowを縦並びに */
1330
+ .row, [class*="row"] {
1331
+ flex-direction: column !important;
1332
+ }
1333
+ }
1334
+
1335
+ /* ========================================
1336
+ アニメーション
1337
+ ======================================== */
1338
+ .tabitem > div {
1339
+ animation: softFadeIn 0.5s ease forwards;
1340
+ }
1341
+ @keyframes softFadeIn {
1342
+ from { opacity: 0; transform: translateY(12px); }
1343
+ to { opacity: 1; transform: translateY(0); }
1344
+ }
1345
+
1346
+ /* ========================================
1347
+ スクロールバー: やわらかく
1348
+ ======================================== */
1349
+ ::-webkit-scrollbar {
1350
+ width: 4px;
1351
+ }
1352
+ ::-webkit-scrollbar-track {
1353
+ background: transparent;
1354
+ }
1355
+ ::-webkit-scrollbar-thumb {
1356
+ background: rgba(200,191,224,0.3);
1357
+ border-radius: 4px;
1358
+ }
1359
+
1360
+ /* ========================================
1361
+ Gradio内部のpadding/border補正
1362
+ ======================================== */
1363
+ .block {
1364
+ border: none !important;
1365
+ background: transparent !important;
1366
+ padding: 0 !important;
1367
+ }
1368
+ .form {
1369
+ background: transparent !important;
1370
+ border: none !important;
1371
+ }
1372
+ .container {
1373
+ background: transparent !important;
1374
+ }
1375
+ .tabs {
1376
+ background: #ffffff !important;
1377
+ }
1378
+
1379
+ /* ========================================
1380
+ 全テキスト強制黒(最終手段)
1381
+ ボタン・ヒーロー系は除外
1382
+ ======================================== */
1383
+ body, body * {
1384
+ color: #1a1a1a !important;
1385
+ }
1386
+
1387
+ /* ヒーローボタン: 色を個別に戻す */
1388
+ .hero-btn--outline,
1389
+ .hero-btn--outline * {
1390
+ color: #555555 !important;
1391
+ }
1392
+ .hero-btn--dark,
1393
+ .hero-btn--dark * {
1394
+ color: #d8d8d8 !important;
1395
+ }
1396
+
1397
+ /* Gradioボタン */
1398
+ button[class*="primary"],
1399
+ button[class*="primary"] * {
1400
+ color: #d8d8d8 !important;
1401
+ }
1402
+ button[class*="secondary"],
1403
+ button[class*="secondary"] * {
1404
+ color: #555555 !important;
1405
+ }
1406
+
1407
+ /* JSON表示の色 */
1408
+ .json-holder *,
1409
+ [data-testid="json"] * {
1410
+ color: #1a1a1a !important;
1411
+ }
1412
+ """
1413
+
1414
+ # ── UI変更: HEAD メタタグ + セリフ体Webフォント読み込み ──
1415
+ HEAD = """
1416
+ <meta name="viewport" content="width=device-width, initial-scale=1, maximum-scale=1, viewport-fit=cover, user-scalable=no">
1417
+ <meta name="theme-color" content="#fdf2f8">
1418
+ <meta name="apple-mobile-web-app-status-bar-style" content="default">
1419
+ <link rel="preconnect" href="https://fonts.googleapis.com">
1420
+ <link rel="preconnect" href="https://fonts.gstatic.com" crossorigin>
1421
+ <link href="https://fonts.googleapis.com/css2?family=Cormorant+Garamond:wght@300;400;500;600&display=swap" rel="stylesheet">
1422
+ """
1423
+
1424
+ # =========================
1425
+ # Gradio UI
1426
+ # ── UI変更: Blocks構造をフェミニン・パステルデザインに再構築 ──
1427
+ # ── 縦スクロール1カラム / 中央寄せ / 余白たっぷり ──
1428
+ # =========================
1429
+ # Gradio 6ではcss/theme/headはmount_gradio_app()で指定
1430
+ with gr.Blocks() as demo:
1431
+
1432
+ # ── UI変更: ヒーローセクションをトップに配置 ──
1433
+ gr.HTML(SENSOR_UI)
1434
+
1435
+ # ── UI変更: セクションタイトル(やわらかいフォント) ──
1436
+ gr.Markdown("## - W O R K F L O W -", elem_classes=["section-title"])
1437
+
1438
+ with gr.Tabs():
1439
+
1440
+ # ── UI変更: 収集タブ — 縦並び1カラム ──
1441
+ with gr.Tab("収集"):
1442
+ label = gr.Dropdown(
1443
+ choices=DEFAULT_LABELS,
1444
+ value=DEFAULT_LABELS[0],
1445
+ label="ラベル"
1446
+ )
1447
+ btn_c_start = gr.Button("収集開始", size="lg")
1448
+ btn_c_stop = gr.Button("収集停止 → 保存", size="lg")
1449
+ collect_msg = gr.Markdown("-")
1450
+ counts_json = gr.JSON(value={"TOTAL": 0}, label="回数カウンタ")
1451
+
1452
+ btn_c_start.click(collect_start, inputs=[label], outputs=[collect_msg, btn_c_start, btn_c_stop])
1453
+ btn_c_stop.click(collect_stop_and_save, inputs=None, outputs=[collect_msg, btn_c_start, btn_c_stop, counts_json])
1454
+
1455
+ # ── UI変更: 学習タブ — スライダー縦並び ──
1456
+ with gr.Tab("学習"):
1457
+ stline = gr.Markdown("MODEL: not trained")
1458
+ window_sec = gr.Slider(0.6, 2.0, value=CFG.window_sec, step=0.1, label="window_sec")
1459
+ hop_sec = gr.Slider(0.1, 0.5, value=CFG.hop_sec, step=0.1, label="hop_sec")
1460
+ fs_target = gr.Slider(20, 100, value=CFG.fs_target, step=5, label="fs_target")
1461
+ feat_mode = gr.Radio(choices=["last", "mean"], value="last", label="state aggregation")
1462
+ btn_train = gr.Button("学習", variant="primary", size="lg")
1463
+ train_msg = gr.Markdown("-")
1464
+ train_cfg = gr.JSON(label="選ばれたハイパラ")
1465
+ tail_log = gr.Textbox(lines=6, label="ログ(末尾)")
1466
+
1467
+ btn_train.click(train_click, inputs=[window_sec, hop_sec, fs_target, feat_mode],
1468
+ outputs=[stline, train_cfg, train_msg, tail_log])
1469
+
1470
+ # ── UI変更: 推論タブ — 予測表示中央 ──
1471
+ with gr.Tab("推論"):
1472
+ infer_state = gr.Markdown("推論: OFF")
1473
+ btn_i_start = gr.Button("推論開始", size="lg")
1474
+ btn_i_stop = gr.Button("推論停止", size="lg")
1475
+ stline2 = gr.Markdown("MODEL: not trained")
1476
+ pred_html = gr.HTML("<div style='font-size:24px;font-weight:800;opacity:.6'>-</div>")
1477
+ prob_json = gr.JSON(label="確信度(クラス別)")
1478
+
1479
+ btn_i_start.click(infer_start, inputs=None, outputs=[infer_state, pred_html, prob_json, stline2])
1480
+ btn_i_stop.click(infer_stop, inputs=None, outputs=[infer_state, pred_html, prob_json, stline2])
1481
+
1482
+ timer_inf = gr.Timer(value=CFG.hop_sec)
1483
+ timer_inf.tick(infer_tick, inputs=None, outputs=[pred_html, prob_json, stline2])
1484
+
1485
+ # ── UI変更: 対話タブ ──
1486
+ with gr.Tab("対話"):
1487
+ gr.Markdown("### 推論結果")
1488
+ chat_big = gr.HTML("<div style='font-size:22px;font-weight:800;opacity:.6'>推論がOFFです</div>")
1489
+ chat_log = gr.Markdown("(log empty)")
1490
+ timer_chat = gr.Timer(value=0.3)
1491
+ timer_chat.tick(chat_tick, inputs=None, outputs=[chat_big, chat_log])
1492
+
1493
+ demo.load(fn=None, inputs=None, outputs=None, js=JS_BOOT)
1494
+
1495
+
1496
+ # =========================
1497
+ # Mount Gradio into FastAPI (SSR OFF)
1498
+ # =========================
1499
+ # ── UI変更: Gradio 6ではcss/theme/head をmount_gradio_appに渡す ──
1500
+ app = gr.mount_gradio_app(
1501
+ api, demo, path="/", ssr_mode=False,
1502
+ css=CSS,
1503
+ head=HEAD,
1504
+ theme=gr.themes.Base(
1505
+ text_size=gr.themes.sizes.text_md,
1506
+ font=["Inter", "Hiragino Kaku Gothic ProN", "Noto Sans JP", "sans-serif"],
1507
+ ).set(
1508
+ body_text_color="#1a1a1a",
1509
+ body_text_color_subdued="#333333",
1510
+ block_label_text_color="#222222",
1511
+ block_title_text_color="#1a1a1a",
1512
+ checkbox_label_text_color="#1a1a1a",
1513
+ table_text_color="#1a1a1a",
1514
+ link_text_color="#333333",
1515
+ color_accent_soft="#e8d5e0",
1516
+ input_background_fill="#ffffff",
1517
+ input_background_fill_dark="#ffffff",
1518
+ input_border_color="#1a1a1a",
1519
+ input_border_color_dark="#1a1a1a",
1520
+ ),
1521
+ )
1522
+
1523
+
1524
+ # =========================
1525
+ # Run on Colab (background thread)
1526
+ # =========================
1527
+ def run_colab(server_port: int = 7860):
1528
+ import uvicorn
1529
+ config = uvicorn.Config(app, host="0.0.0.0", port=server_port, log_level="warning")
1530
+ server = uvicorn.Server(config)
1531
+
1532
+ th = threading.Thread(target=server.run, daemon=True)
1533
+ th.start()
1534
+ time.sleep(1.0)
1535
+
1536
+ # If in colab, show public URL via Gradio share too (simpler UX)
1537
+ # Note: We can't "launch" gradio separately because FastAPI+mount is already serving.
1538
+ # So: use Colab's port proxy link if available, otherwise open localhost in browser.
1539
+ try:
1540
+ from google.colab import output
1541
+ proxy_url = output.eval_js(f"google.colab.kernel.proxyPort({server_port})")
1542
+ print("Open this URL (PC):", proxy_url)
1543
+ print("Open the same URL on your smartphone to use accelerometer.")
1544
+ except Exception:
1545
+ print(f"Server running on http://127.0.0.1:{server_port} (Colab proxy unavailable here).")
1546
+
1547
+ # Start
1548
+ if "google.colab" in sys.modules:
1549
+ run_colab(7860)
1550
+ else:
1551
+ # local python run (non-notebook)
1552
+ import uvicorn
1553
+ uvicorn.run(app, host="0.0.0.0", port=7860)