IDKHowToCodeFr/tinyml-backend
1
1import os2import joblib3import json4import numpy as np5from sklearn.neighbors import KNeighborsClassifier6from sklearn.svm import SVC7from sklearn.linear_model import LogisticRegression8from sklearn.ensemble import RandomForestClassifier9from sklearn.neural_network import MLPClassifier10from sklearn.metrics import f1_score11import pandas as pd12import sys13 14sys.path.append(os.path.dirname(os.path.abspath(__file__)))15from preprocessing import get_train_test_split, resolve_model_dir16 17def train_models():18 data_path = os.path.join(os.path.dirname(os.path.abspath(__file__)), '..', 'data', 'patient_dataset.csv')19 model_dir = resolve_model_dir()20 os.makedirs(model_dir, exist_ok=True)21 registry_path = os.path.join(model_dir, 'registry.json')22 23 if os.path.exists(registry_path):24 with open(registry_path, 'r') as f:25 registry = json.load(f)26 else:27 registry = {"active_version": "v1", "versions": {"v1": {"score": 0.0}}}28 29 print("Loading data...")30 df = pd.read_csv(data_path, encoding='utf-8')31 X_train, X_test, y_train, y_test = get_train_test_split(df)32 33 models = {34 'knn': KNeighborsClassifier(n_neighbors=5),35 'svm': SVC(kernel='linear', probability=True, max_iter=2000), 36 'logreg': LogisticRegression(max_iter=1000),37 'rf': RandomForestClassifier(n_estimators=30, max_depth=5, random_state=42),38 'small_nn': MLPClassifier(hidden_layer_sizes=(16, 8), max_iter=500, random_state=42)39 }40 41 trained_models = {}42 scores = []43 44 for name, model in models.items():45 print(f"Training {name}...")46 model.fit(X_train, y_train)47 trained_models[name] = model48 49 preds = model.predict(X_test)50 f1 = f1_score(y_test, preds, average='weighted')51 scores.append(f1)52 53 avg_score = sum(scores) / len(scores)54 new_version_num = len(registry["versions"]) + 155 new_version = f"v{new_version_num}"56 print(f"New Version {new_version} Average F1 Score: {avg_score:.4f}")57 58 registry["versions"][new_version] = {"score": avg_score}59 active_version = registry["active_version"]60 active_score = registry["versions"].get(active_version, {}).get("score", 0.0)61 62 if avg_score >= active_score:63 print(f"Score improved ({avg_score:.4f} >= {active_score:.4f}). Overwriting active models.")64 for name, model in trained_models.items():65 joblib.dump(model, os.path.join(model_dir, f'{name}.pkl'))66 registry["active_version"] = new_version67 else:68 print(f"Score degraded ({avg_score:.4f} < {active_score:.4f}). Rollback initiated: keeping existing models.")69 70 with open(registry_path, 'w') as f:71 json.dump(registry, f, indent=4)72 73 print("Training Complete.")74 75if __name__ == '__main__':76 train_models()77 