Rthur2003 commited on
Commit
4324aed
·
1 Parent(s): bb33480

feat: enhance model compatibility and input handling for deep learning wrappers

Browse files
Files changed (1) hide show
  1. local_demo.py +35 -7
local_demo.py CHANGED
@@ -79,6 +79,7 @@ def _safe_model_name(name: str) -> str:
79
  .replace("(", "")
80
  .replace(")", "")
81
  .replace("/", "_")
 
82
  )
83
 
84
 
@@ -132,10 +133,16 @@ def _load_dataset_summary() -> dict[str, Any]:
132
 
133
 
134
  def _find_matching_name(raw_name: str, training_results: dict[str, Any]) -> str:
 
 
 
 
 
135
  for name in training_results:
136
  if name.startswith("_"):
137
  continue
138
- if _safe_model_name(name) == raw_name:
 
139
  return name
140
  return raw_name.replace("_", " ").title()
141
 
@@ -157,6 +164,15 @@ def _load_artifacts() -> DemoArtifacts:
157
  scaler = _load_pickle(scaler_path)
158
  feature_cols = _load_json(columns_path)
159
  training_results = _load_json(results_path)
 
 
 
 
 
 
 
 
 
160
  feature_importance = training_results.get("_feature_importance", {})
161
  best_model_name = training_results.get("_best_model", "Gradient Boosting")
162
 
@@ -327,6 +343,11 @@ def _build_signal_html(features: dict[str, float]) -> str:
327
  return "".join(parts)
328
 
329
 
 
 
 
 
 
330
  def _build_model_table_html(
331
  selected_model_name: str,
332
  feature_vector: np.ndarray,
@@ -338,7 +359,9 @@ def _build_model_table_html(
338
  try:
339
  with warnings.catch_warnings():
340
  warnings.simplefilter("ignore", category=UserWarning)
341
- probability = float(model.predict_proba(scaled)[0][1])
 
 
342
  except Exception: # noqa: BLE001
343
  continue
344
  scored_rows.append((name, probability))
@@ -461,11 +484,15 @@ def analyze_audio(audio_file: Any, selected_model_label: str) -> tuple[str, str,
461
  return error_html, "", "", ""
462
 
463
  feature_vector = _build_feature_vector(features)
464
- scaled = ARTIFACTS.scaler.transform(feature_vector)
465
  model = ARTIFACTS.loaded_models[selected_model_name]
 
 
 
 
 
466
  with warnings.catch_warnings():
467
  warnings.simplefilter("ignore", category=UserWarning)
468
- ai_prob = float(model.predict_proba(scaled)[0][1])
469
  elapsed = time.time() - start_time
470
 
471
  result_html = _build_result_html(ai_prob, duration, elapsed, selected_model_name)
@@ -485,8 +512,8 @@ def build_models_md() -> str:
485
  f"- Ozellik sayisi: **{training_results.get('_n_features', len(ARTIFACTS.feature_cols))}**",
486
  f"- CV kat sayisi: **{training_results.get('_n_folds', '?')}**",
487
  "",
488
- "| Model | CV AUC | Holdout AUC | Acc | F1 |",
489
- "|------|--------|-------------|-----|----|",
490
  ]
491
 
492
  model_names = [
@@ -499,8 +526,9 @@ def build_models_md() -> str:
499
  for name in model_names:
500
  result = training_results[name]
501
  display = f"**{name}**" if name == ARTIFACTS.best_model_name else name
 
502
  lines.append(
503
- f"| {display} | {result.get('roc_auc', 0.0):.4f} | "
504
  f"{result.get('validation_auc', 0.0):.4f} | "
505
  f"{result.get('accuracy', 0.0):.4f} | {result.get('f1', 0.0):.4f} |"
506
  )
 
79
  .replace("(", "")
80
  .replace(")", "")
81
  .replace("/", "_")
82
+ .replace("-", "_")
83
  )
84
 
85
 
 
133
 
134
 
135
  def _find_matching_name(raw_name: str, training_results: dict[str, Any]) -> str:
136
+ # DL model files are stored as "model_dl_<name>" — strip the "dl_" infix too
137
+ candidates = [raw_name]
138
+ if raw_name.startswith("dl_"):
139
+ candidates.append(raw_name[3:]) # "dl_deep_mlp_..." -> "deep_mlp_..."
140
+
141
  for name in training_results:
142
  if name.startswith("_"):
143
  continue
144
+ safe = _safe_model_name(name)
145
+ if safe in candidates:
146
  return name
147
  return raw_name.replace("_", " ").title()
148
 
 
164
  scaler = _load_pickle(scaler_path)
165
  feature_cols = _load_json(columns_path)
166
  training_results = _load_json(results_path)
167
+
168
+ # Merge DL metrics into the unified training_results dict
169
+ dl_results_path = MODELS_DIR / "deep_learning_results.json"
170
+ if dl_results_path.exists():
171
+ dl_results = _load_json(dl_results_path)
172
+ for name, metrics in dl_results.items():
173
+ if name not in training_results:
174
+ training_results[name] = metrics
175
+
176
  feature_importance = training_results.get("_feature_importance", {})
177
  best_model_name = training_results.get("_best_model", "Gradient Boosting")
178
 
 
343
  return "".join(parts)
344
 
345
 
346
+ def _is_dl_wrapper(model: Any) -> bool:
347
+ """True for TorchSklearnWrapper — it has its own internal scaler."""
348
+ return type(model).__name__ == "TorchSklearnWrapper"
349
+
350
+
351
  def _build_model_table_html(
352
  selected_model_name: str,
353
  feature_vector: np.ndarray,
 
359
  try:
360
  with warnings.catch_warnings():
361
  warnings.simplefilter("ignore", category=UserWarning)
362
+ # DL wrappers scale internally — pass raw vector to avoid double-scaling
363
+ input_vector = feature_vector if _is_dl_wrapper(model) else scaled
364
+ probability = float(model.predict_proba(input_vector)[0][1])
365
  except Exception: # noqa: BLE001
366
  continue
367
  scored_rows.append((name, probability))
 
484
  return error_html, "", "", ""
485
 
486
  feature_vector = _build_feature_vector(features)
 
487
  model = ARTIFACTS.loaded_models[selected_model_name]
488
+ # DL wrappers scale internally; ML models expect pre-scaled input
489
+ if _is_dl_wrapper(model):
490
+ input_vector = feature_vector
491
+ else:
492
+ input_vector = ARTIFACTS.scaler.transform(feature_vector)
493
  with warnings.catch_warnings():
494
  warnings.simplefilter("ignore", category=UserWarning)
495
+ ai_prob = float(model.predict_proba(input_vector)[0][1])
496
  elapsed = time.time() - start_time
497
 
498
  result_html = _build_result_html(ai_prob, duration, elapsed, selected_model_name)
 
512
  f"- Ozellik sayisi: **{training_results.get('_n_features', len(ARTIFACTS.feature_cols))}**",
513
  f"- CV kat sayisi: **{training_results.get('_n_folds', '?')}**",
514
  "",
515
+ "| Model | Tip | CV AUC | Holdout AUC | Acc | F1 |",
516
+ "|------|-----|--------|-------------|-----|----|",
517
  ]
518
 
519
  model_names = [
 
526
  for name in model_names:
527
  result = training_results[name]
528
  display = f"**{name}**" if name == ARTIFACTS.best_model_name else name
529
+ model_type = "DL" if result.get("type") == "deep_learning" else "ML"
530
  lines.append(
531
+ f"| {display} | {model_type} | {result.get('roc_auc', 0.0):.4f} | "
532
  f"{result.get('validation_auc', 0.0):.4f} | "
533
  f"{result.get('accuracy', 0.0):.4f} | {result.get('f1', 0.0):.4f} |"
534
  )