Spaces:
Running
Running
feat: enhance model compatibility and input handling for deep learning wrappers
Browse files- 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 |
-
|
|
|
|
| 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 |
-
|
|
|
|
|
|
|
| 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(
|
| 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 |
)
|