Text Classification
Keras
English
Azerbaijani
prompt-injection
security
llm-security
document-security
retvec
cnn
tensorflow
fastapi
Eval Results (legacy)
Instructions to use MegrurNiftiyev/MyGuard-Prompt-Injection-Detector with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Keras
How to use MegrurNiftiyev/MyGuard-Prompt-Injection-Detector with Keras:
# Available backend options are: "jax", "torch", "tensorflow". import os os.environ["KERAS_BACKEND"] = "jax" import keras model = keras.saving.load_model("hf://MegrurNiftiyev/MyGuard-Prompt-Injection-Detector") - Notebooks
- Google Colab
- Kaggle
| """ | |
| Standalone RETVec+CNN Keras model training & held-out test evaluation script. | |
| Usage: | |
| python train_model.py | |
| python -m app.scripts.train_model | |
| """ | |
| import os | |
| import sys | |
| import random | |
| import zipfile | |
| import docx | |
| import pypdf | |
| from pptx import Presentation | |
| import numpy as np | |
| os.environ["TF_USE_LEGACY_KERAS"] = "1" | |
| os.environ["CUDA_VISIBLE_DEVICES"] = "-1" | |
| sys.stdout.reconfigure(encoding='utf-8') | |
| # Ensure project root is in sys.path | |
| BASE_DIR = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..")) | |
| if BASE_DIR not in sys.path: | |
| sys.path.insert(0, BASE_DIR) | |
| SEED = 42 | |
| random.seed(SEED) | |
| np.random.seed(SEED) | |
| import tensorflow as tf | |
| tf.random.set_seed(SEED) | |
| from app.ml.cnn.architecture import build_model, LABEL_NAMES | |
| from app.ml.training.data.encoding import encode_labels | |
| from app.ml.training.train import get_class_weights | |
| from app.ml.preprocessing.chunking import chunk_text | |
| HELDOUT_TEST_FILES = { | |
| "benign": [ | |
| "09_resmi_mektub_temiz.docx", | |
| "10_iclas_protokolu_temiz.docx", | |
| "Monthly Financial Expense Report.pdf", | |
| "11_ezamiyye_emri_temiz.docx", | |
| "19_sifaris_senedi_temiz.docx" | |
| ], | |
| "injection": [ | |
| "01_Aylıq_Fəaliyyət_Hesabatı.docx", | |
| "16_ezamiyye_xercleri_injection_gizli.docx", | |
| "19_sifaris_senedi_problem.docx", | |
| "23_bank_zemanet_mektubu_injection_context_hijack.docx", | |
| "24_qebul_tehvil_akti_injection.docx" | |
| ] | |
| } | |
| def extract_pptx(file_path: str) -> str: | |
| """Extract slide paragraph text and notes text from PPTX files using python-pptx.""" | |
| try: | |
| prs = Presentation(file_path) | |
| parts = [] | |
| for slide in prs.slides: | |
| for shape in slide.shapes: | |
| if shape.has_text_frame: | |
| for para in shape.text_frame.paragraphs: | |
| line = "".join(run.text for run in para.runs) | |
| if line.strip(): | |
| parts.append(line.strip()) | |
| if slide.has_notes_slide and slide.notes_slide.notes_text_frame: | |
| note = slide.notes_slide.notes_text_frame.text | |
| if note.strip(): | |
| parts.append(note.strip()) | |
| return "\n".join(parts) | |
| except Exception as e: | |
| print(f"Warning reading PPTX {file_path}: {e}") | |
| return "" | |
| def extract_text(file_path: str) -> str: | |
| """Extract raw text from supported document formats (.docx, .pptx, .pdf, .zip, .txt).""" | |
| ext = os.path.splitext(file_path)[1].lower() | |
| text = "" | |
| try: | |
| if ext == ".docx": | |
| doc = docx.Document(file_path) | |
| parts = [p.text for p in doc.paragraphs if p.text.strip()] | |
| for table in doc.tables: | |
| for row in table.rows: | |
| for cell in row.cells: | |
| if cell.text.strip(): | |
| parts.append(cell.text.strip()) | |
| text = "\n".join(parts) | |
| elif ext == ".pptx": | |
| text = extract_pptx(file_path) | |
| elif ext == ".pdf": | |
| reader = pypdf.PdfReader(file_path) | |
| parts = [] | |
| for i, page in enumerate(reader.pages): | |
| if i >= 20: | |
| break | |
| try: | |
| t = page.extract_text() | |
| if t: | |
| parts.append(t.strip()) | |
| except Exception: | |
| continue | |
| text = "\n".join(parts) | |
| elif ext == ".zip": | |
| parts = [] | |
| with zipfile.ZipFile(file_path, 'r') as z: | |
| for name in z.namelist(): | |
| if name.endswith('.docx'): | |
| tmp_path = os.path.join(os.path.dirname(file_path), "_tmp_extracted.docx") | |
| with open(tmp_path, "wb") as f_out: | |
| f_out.write(z.read(name)) | |
| sub_text = extract_text(tmp_path) | |
| if os.path.exists(tmp_path): | |
| os.remove(tmp_path) | |
| parts.append(sub_text) | |
| elif name.endswith('.pptx'): | |
| tmp_path = os.path.join(os.path.dirname(file_path), "_tmp_extracted.pptx") | |
| with open(tmp_path, "wb") as f_out: | |
| f_out.write(z.read(name)) | |
| sub_text = extract_text(tmp_path) | |
| if os.path.exists(tmp_path): | |
| os.remove(tmp_path) | |
| parts.append(sub_text) | |
| elif name.endswith('.txt'): | |
| parts.append(z.read(name).decode('utf-8', errors='ignore')) | |
| text = "\n".join(parts) | |
| elif ext == ".txt": | |
| with open(file_path, "r", encoding="utf-8", errors="ignore") as f: | |
| text = f.read() | |
| else: | |
| print(f"Skipping unsupported file extension {ext} for {file_path}") | |
| return "" | |
| except Exception as e: | |
| print(f"Warning reading {file_path}: {e}") | |
| return text.strip() | |
| def split_documents(doc_ids: list[str], val_ratio: float = 0.15, seed: int = 42) -> tuple[set[str], set[str]]: | |
| """Perform a document-level split of source document IDs into train and validation sets.""" | |
| rng = random.Random(seed) | |
| unique_ids = list(dict.fromkeys(doc_ids)) | |
| rng.shuffle(unique_ids) | |
| n_val = max(1, int(len(unique_ids) * val_ratio)) | |
| val_ids = set(unique_ids[:n_val]) | |
| train_ids = set(unique_ids[n_val:]) | |
| return train_ids, val_ids | |
| def load_real_dataset(raw_dir: str): | |
| all_chunks = [] # [(doc_id, text_chunk, label)] | |
| all_doc_ids = [] | |
| test_docs = [] | |
| # Define folder mapping: (folder_path, default_category) | |
| folders_to_scan = [ | |
| (os.path.join(raw_dir, "benign"), "benign"), | |
| (os.path.join(raw_dir, "injection"), "injection"), | |
| ] | |
| downloaded_dir = os.path.join(raw_dir, "downloaded") | |
| if os.path.exists(downloaded_dir): | |
| for root, dirs, files in os.walk(downloaded_dir): | |
| if files: | |
| folders_to_scan.append((root, "benign")) | |
| scanned_file_counts = {} | |
| # Load 10,200 PDF V4 Synthetic Dataset if dataset_V4.csv exists | |
| v4_csv_path = os.path.join(downloaded_dir, "dataset_V4.csv") | |
| if os.path.exists(v4_csv_path): | |
| try: | |
| import pandas as pd | |
| print(f"Loading 10,200 PDF V4 Synthetic Dataset samples from {v4_csv_path}...") | |
| df_v4 = pd.read_csv(v4_csv_path) | |
| v4_count = 0 | |
| for _, row in df_v4.iterrows(): | |
| doc_id = f"v4_{row['doc_id']}" | |
| extracted_text = str(row['extracted_text']) if pd.notna(row['extracted_text']) else "" | |
| if not extracted_text.strip(): | |
| continue | |
| is_inj = bool(row['is_injected']) | |
| lbl = "injection" if is_inj else "safe" | |
| v4_count += 1 | |
| lines = [l.strip() for l in extracted_text.split("\n") if l.strip()] | |
| for line in lines: | |
| words = line.split() | |
| if len(words) <= 60: | |
| all_chunks.append((doc_id, line, lbl)) | |
| all_doc_ids.append(doc_id) | |
| else: | |
| for c in chunk_text(line): | |
| all_chunks.append((doc_id, c, lbl)) | |
| all_doc_ids.append(doc_id) | |
| scanned_file_counts["dataset_V4.csv (10,200 PDFs)"] = v4_count | |
| except Exception as err: | |
| print(f"Warning loading dataset_V4.csv: {err}") | |
| for cat_dir, category in folders_to_scan: | |
| if not os.path.exists(cat_dir): | |
| continue | |
| heldout_list = HELDOUT_TEST_FILES.get(category, []) | |
| label_str = "safe" if category == "benign" else "injection" | |
| dir_key = os.path.relpath(cat_dir, raw_dir) | |
| scanned_file_counts[dir_key] = scanned_file_counts.get(dir_key, 0) | |
| for fname in os.listdir(cat_dir): | |
| fpath = os.path.join(cat_dir, fname) | |
| if not os.path.isfile(fpath): | |
| continue | |
| extracted = extract_text(fpath) | |
| if not extracted: | |
| continue | |
| scanned_file_counts[dir_key] += 1 | |
| doc_id = os.path.relpath(fpath, raw_dir) | |
| if fname in heldout_list: | |
| test_docs.append({ | |
| "filename": fname, | |
| "category": category, | |
| "expected_label": label_str, | |
| "text": extracted | |
| }) | |
| else: | |
| ext = os.path.splitext(fname)[1].lower() | |
| file_chunks = [] | |
| # Check for docx paragraph-level white font / hidden text | |
| docx_inj_lines = set() | |
| if ext == ".docx": | |
| try: | |
| doc = docx.Document(fpath) | |
| for p in doc.paragraphs: | |
| ptxt = p.text.strip() | |
| if not ptxt: | |
| continue | |
| is_p_white = False | |
| for r in p.runs: | |
| if r.font.color and r.font.color.rgb and str(r.font.color.rgb).upper() in ("FFFFFF", "FFF"): | |
| is_p_white = True | |
| break | |
| if r._r.rPr is not None and r._r.rPr.find(docx.oxml.ns.qn("w:vanish")) is not None: | |
| is_p_white = True | |
| break | |
| if is_p_white: | |
| docx_inj_lines.add(ptxt) | |
| except Exception: | |
| pass | |
| MANUAL_STEALTHY_LABELS = { | |
| "23_bank_zemanet_mektubu_injection_context_hijack.docx": "450,000 azn", | |
| "24_qebul_tehvil_akti_injection.docx": "administrative correction", | |
| "25_sigorta_polisi_injection.docx": "franchise override protocol", | |
| "26_emek_muqavilesi_inject.docx": "critical parser override", | |
| "27_vekaletname_inject.docx": "müvəqqəti hüquqi düzəliş", | |
| "29_bank_rekvizit_deyisikligi_injection.docx": "sistem təlimatı", | |
| "31_tecili_odenis_memo_inject.docx": "suni intelleke", | |
| "32_hosting_deyisiklik_bildirisi_inject.docx": "audit id-7741", | |
| "33_elave_is_sifarisi_inject.docx": "audit id-112", | |
| "34_distributor_endirim_teklifi_inject.docx": "audit id-267", | |
| "38_icare_muqavilesi_injection.docx": "tərəfdaşlıq ianəsi", | |
| "39_dasima_xidmeti_muqavilesi_inject.docx": "<last update>" | |
| } | |
| lines = [l.strip() for l in extracted.split("\n") if l.strip()] | |
| for line in lines: | |
| is_inj_line = False | |
| if category == "injection": | |
| low = line.lower() | |
| if fname in MANUAL_STEALTHY_LABELS: | |
| if MANUAL_STEALTHY_LABELS[fname] in low: | |
| is_inj_line = True | |
| else: | |
| low = line.lower() | |
| if line in docx_inj_lines or any(kw in low for kw in [ | |
| "prompt", "system", "yuxarida", "mene", "ignore", "override", | |
| "@", "//", "#", "||", "^^", "***", "&&", "<system", "[system", | |
| "internal system update", "forget", "unrestricted" | |
| ]): | |
| is_inj_line = True | |
| lbl = "injection" if (category == "injection" and is_inj_line) else "safe" | |
| words = line.split() | |
| if len(words) <= 60: | |
| file_chunks.append((line, lbl)) | |
| else: | |
| for c in chunk_text(line): | |
| file_chunks.append((c, lbl)) | |
| # Cap per-document safe chunks so long PDFs don't dominate dataset (max 15 safe chunks per doc) | |
| if len(file_chunks) > 15: | |
| inj_chunks = [c for c in file_chunks if c[1] == "injection"] | |
| safe_chunks = [c for c in file_chunks if c[1] == "safe"] | |
| needed_safe = max(5, 15 - len(inj_chunks)) | |
| step = max(1, len(safe_chunks) // needed_safe) if safe_chunks else 1 | |
| file_chunks = inj_chunks + (safe_chunks[::step][:needed_safe] if safe_chunks else []) | |
| for text_chunk, lbl in file_chunks: | |
| all_chunks.append((doc_id, text_chunk, lbl)) | |
| all_doc_ids.append(doc_id) | |
| # Document-level split | |
| train_doc_ids, val_doc_ids = split_documents(all_doc_ids, val_ratio=0.15, seed=SEED) | |
| train_tuples = [c for c in all_chunks if c[0] in train_doc_ids] | |
| val_tuples = [c for c in all_chunks if c[0] in val_doc_ids] | |
| # Oversample injection training tuples so model learns injection patterns properly | |
| train_inj_tuples = [t for t in train_tuples if t[2] == "injection"] | |
| train_safe_tuples = [t for t in train_tuples if t[2] == "safe"] | |
| if train_inj_tuples and len(train_safe_tuples) > 0: | |
| multiplier = max(1, (len(train_safe_tuples) // 3) // len(train_inj_tuples)) | |
| train_inj_oversampled = train_inj_tuples * multiplier | |
| train_tuples = train_safe_tuples + train_inj_oversampled | |
| # Thorough random shuffling across all sources, classes, and languages | |
| rng = random.Random(SEED) | |
| rng.shuffle(train_tuples) | |
| rng.shuffle(val_tuples) | |
| train_texts = [t[1] for t in train_tuples] | |
| train_labels = [t[2] for t in train_tuples] | |
| val_texts = [t[1] for t in val_tuples] | |
| val_labels = [t[2] for t in val_tuples] | |
| print("Scanned files count per folder:") | |
| for folder_rel, count in scanned_file_counts.items(): | |
| print(f" - {folder_rel}: {count} valid documents") | |
| print(f"Document-level split: {len(train_doc_ids)} train docs ({len(train_texts)} chunks), {len(val_doc_ids)} val docs ({len(val_texts)} chunks)") | |
| return (train_texts, train_labels), (val_texts, val_labels), test_docs | |
| def main(): | |
| raw_dir = os.path.join(BASE_DIR, "data", "raw") | |
| print("Reading document dataset from data/raw...") | |
| (train_texts, train_labels), (val_texts, val_labels), test_docs = load_real_dataset(raw_dir) | |
| print(f"\n--- Dataset Loading Summary ---") | |
| print(f"Training text chunks extracted: {len(train_texts)}") | |
| print(f" - Safe (Benign) train chunks: {train_labels.count('safe')}") | |
| print(f" - Injection train chunks: {train_labels.count('injection')}") | |
| print(f"Validation text chunks extracted: {len(val_texts)}") | |
| print(f" - Safe (Benign) val chunks: {val_labels.count('safe')}") | |
| print(f" - Injection val chunks: {val_labels.count('injection')}") | |
| print(f"Held-out Test Files reserved: {len(test_docs)}") | |
| for td in test_docs: | |
| print(f" * [{td['category'].upper()}] {td['filename']} ({len(td['text'])} chars)") | |
| X_train = np.array([[t] for t in train_texts]) | |
| Y_train_label = encode_labels(train_labels) | |
| X_val = np.array([[t] for t in val_texts]) | |
| Y_val_label = encode_labels(val_labels) | |
| class_weights_dict = get_class_weights(Y_train_label) | |
| sample_weights_label = np.array([class_weights_dict[int(np.argmax(y))] for y in Y_train_label], dtype=np.float32) | |
| print("\nBuilding RETVec + CNN Keras Classification Model...") | |
| model = build_model(sequence_length=128) | |
| model.summary() | |
| print("\nStarting Keras Model Training (5 Epochs, batch_size=128, document-level validation)...", flush=True) | |
| history = model.fit( | |
| X_train, | |
| Y_train_label, | |
| epochs=5, | |
| batch_size=128, | |
| validation_data=(X_val, Y_val_label), | |
| sample_weight=sample_weights_label, | |
| verbose=1 | |
| ) | |
| models_dir = os.path.join(BASE_DIR, "data", "models") | |
| os.makedirs(models_dir, exist_ok=True) | |
| keras_model_path = os.path.join(models_dir, "retvec_cnn_model.keras") | |
| print(f"\nSaving trained model to .keras file at:\n {keras_model_path}") | |
| model.save(keras_model_path) | |
| cache_dir = os.path.join(BASE_DIR, "data", "cache") | |
| os.makedirs(cache_dir, exist_ok=True) | |
| model.save(os.path.join(cache_dir, "active_model.keras")) | |
| print("\n==========================================") | |
| print("HELD-OUT TEST FILES INFERENCE & EVALUATION") | |
| print("==========================================") | |
| correct_predictions = 0 | |
| test_results = [] | |
| for td in test_docs: | |
| raw_text = td["text"] | |
| lines = [l.strip() for l in raw_text.split("\n") if l.strip()] | |
| chunks = [] | |
| for line in lines: | |
| words = line.split() | |
| if len(words) <= 60: | |
| chunks.append(line) | |
| else: | |
| chunks.extend(chunk_text(line)) | |
| chunk_inputs = np.array([[c] for c in chunks]) | |
| preds = model.predict(chunk_inputs, verbose=0) | |
| label_preds = preds if isinstance(preds, np.ndarray) and preds.ndim == 2 else preds[0] | |
| worst_chunk_idx = label_preds[:, 2].argmax() | |
| max_injection_prob = float(label_preds[worst_chunk_idx, 2]) | |
| max_inj_line = chunks[worst_chunk_idx] if chunks else "" | |
| avg_probs = np.mean(label_preds, axis=0) | |
| HIGH_CONF_THRESHOLD = 0.85 | |
| CORROBORATION_THRESHOLD = 0.60 | |
| MIN_CORROBORATING_CHUNKS = 2 | |
| injection_probs = [float(p) for p in label_preds[:, 2]] | |
| predicted_label = "safe" | |
| high_conf = [p for p in injection_probs if p >= HIGH_CONF_THRESHOLD] | |
| if high_conf: | |
| predicted_label = "injection" | |
| else: | |
| corroborating = [p for p in injection_probs if p >= CORROBORATION_THRESHOLD] | |
| if len(corroborating) >= MIN_CORROBORATING_CHUNKS: | |
| predicted_label = "injection" | |
| is_correct = (predicted_label == td["expected_label"]) | |
| if is_correct: | |
| correct_predictions += 1 | |
| test_results.append({ | |
| "filename": td["filename"], | |
| "expected": td["expected_label"], | |
| "predicted": predicted_label, | |
| "is_correct": is_correct, | |
| "prob_safe": float(avg_probs[0]), | |
| "prob_suspicious": float(avg_probs[1]), | |
| "prob_injection": float(avg_probs[2]), | |
| "max_chunk_injection": float(max_injection_prob), | |
| "max_inj_snippet": max_inj_line[:60] | |
| }) | |
| status = "PASSED ✓" if is_correct else "FAILED ✗" | |
| print(f"File: {td['filename']}") | |
| print(f" Expected: {td['expected_label']} | Predicted: {predicted_label} [{status}]") | |
| print(f" Max Injection Prob: {max_injection_prob:.2%} | Snippet: {max_inj_line[:70]!r}\n") | |
| accuracy = (correct_predictions / len(test_docs)) * 100 if test_docs else 0.0 | |
| print(f"Final Held-Out Test Accuracy: {accuracy:.2f}% ({correct_predictions}/{len(test_docs)})") | |
| last_loss = float(history.history["loss"][-1]) if "history" in locals() and "loss" in history.history else 0.0 | |
| last_acc = float(history.history["accuracy"][-1]) if "history" in locals() and "accuracy" in history.history else 0.0 | |
| last_val = float(history.history["val_accuracy"][-1]) if "history" in locals() and "val_accuracy" in history.history else 0.0 | |
| prompt_local_push_confirmation( | |
| model=model, | |
| accuracy=accuracy, | |
| correct_count=correct_predictions, | |
| total_test_docs=len(test_docs), | |
| train_chunk_count=len(train_texts), | |
| val_chunk_count=len(val_texts), | |
| last_train_loss=last_loss, | |
| last_train_acc=last_acc, | |
| last_val_acc=last_val, | |
| test_results=test_results, | |
| ) | |
| def fetch_last_5_models_from_firestore(): | |
| init_firebase() | |
| db = get_firestore_db() | |
| if db is None: | |
| return [], 0 | |
| try: | |
| docs = db.collection("models").get() | |
| model_list = [] | |
| max_run_num = 0 | |
| for doc in docs: | |
| d = doc.to_dict() | |
| v_id = d.get("version") or doc.id | |
| if v_id.startswith("run-"): | |
| try: | |
| r_num = int(v_id.split("-")[1]) | |
| if r_num > max_run_num: | |
| max_run_num = r_num | |
| except ValueError: | |
| pass | |
| model_list.append(d) | |
| def sort_key(d): | |
| v = d.get("version", "") | |
| if v.startswith("run-"): | |
| try: | |
| return int(v.split("-")[1]) | |
| except ValueError: | |
| pass | |
| return 0 | |
| model_list.sort(key=sort_key) | |
| return model_list[-5:], max_run_num | |
| except Exception as e: | |
| print(f"Warning fetching models from Firestore: {e}") | |
| return [], 0 | |
| def prompt_local_push_confirmation(model, accuracy: float, correct_count: int, total_test_docs: int, train_chunk_count: int, val_chunk_count: int, last_train_loss: float, last_train_acc: float, last_val_acc: float, test_results: list): | |
| import subprocess | |
| import asyncio | |
| from datetime import datetime, timezone | |
| from app.core.firebase import init_firebase, get_firestore_db | |
| from app.ml.serving.registry import save_model_version | |
| last_5, max_run_num = fetch_last_5_models_from_firestore() | |
| inj_docs = [t for t in test_results if t["expected"] == "injection"] | |
| inj_correct = [t for t in inj_docs if t["is_correct"]] | |
| test_recall = (len(inj_correct) / len(inj_docs) * 100.0) if inj_docs else 100.0 | |
| if last_5: | |
| print("\nLast 5 registered versions:") | |
| for m in last_5: | |
| v_str = m.get("version", "unknown") | |
| metrics_m = m.get("metrics", {}) | |
| test_acc_m = metrics_m.get("test_acc", 0.0) * 100.0 if isinstance(metrics_m.get("test_acc"), (int, float)) else 0.0 | |
| recall_m = metrics_m.get("recall", 0.0) * 100.0 if isinstance(metrics_m.get("recall"), (int, float)) else 0.0 | |
| status_tag = " (currently active)" if m.get("status") == "active" else "" | |
| print(f" {v_str:<8} test acc {test_acc_m:.2f}% recall {recall_m:.0f}%{status_tag}") | |
| print(f"\nThis run: test acc {accuracy:.2f}% recall {test_recall:.0f}%\n") | |
| answer = input("Upload this model to Firebase as a new candidate version? (y/n): ").strip().lower() | |
| if answer == "y": | |
| next_run_num = max_run_num + 1 if max_run_num > 0 else 12 | |
| new_version_id = f"run-{next_run_num:02d}" | |
| try: | |
| res = subprocess.run(["git", "rev-parse", "HEAD"], capture_output=True, text=True, check=True) | |
| source_commit = res.stdout.strip() | |
| except Exception: | |
| source_commit = "unknown" | |
| today_str = datetime.now(timezone.utc).strftime("%Y-%m-%d") | |
| desc = ( | |
| f"Trained {today_str}. " | |
| f"Dataset: {train_chunk_count} train chunks + {val_chunk_count} val chunks. " | |
| f"Held-out test: {accuracy:.2f}% accuracy ({correct_count}/{total_test_docs}), " | |
| f"{test_recall:.0f}% injection recall." | |
| ) | |
| metrics_payload = { | |
| "train_loss": float(last_train_loss), | |
| "train_acc": float(last_train_acc), | |
| "val_acc": float(last_val_acc), | |
| "test_acc": float(accuracy / 100.0), | |
| "recall": float(test_recall / 100.0), | |
| "correct_test": f"{correct_count}/{total_test_docs}", | |
| } | |
| asyncio.run( | |
| save_model_version( | |
| model=model, | |
| metrics=metrics_payload, | |
| version=new_version_id, | |
| status="candidate", | |
| source_commit=source_commit, | |
| description=desc, | |
| ) | |
| ) | |
| print(f"Uploaded as candidate version '{new_version_id}'. Use POST /model/change-version/{new_version_id} to make it active.") | |
| else: | |
| print("Skipped. Model saved locally only at data/models/retvec_cnn_model.keras.") | |
| if __name__ == "__main__": | |
| main() | |