pvyas96/random_label_forest
0
1import streamlit as st2import numpy as np3from scipy.sparse import csr_matrix4import joblib5from sklearn.metrics import f1_score6from sklearn.decomposition import TruncatedSVD7from tqdm import tqdm8import tempfile9import os10 11st.set_page_config(page_title="RLF Model Tester", layout="wide")12st.title("Random Label Forest (RLF) - Test Dataset Evaluation")13 14# ----------------------- Parsing Function ----------------------- #15def parse_rcv1_sparse_file(filename):16 with open(filename, 'r') as f:17 header = f.readline()18 num_samples, num_features, num_labels = map(int, header.strip().split())19 20 feat_data, feat_rows, feat_cols = [], [], []21 label_data, label_rows, label_cols = [], [], []22 23 for row_idx, line in enumerate(f):24 line = line.strip()25 if not line:26 continue27 label_part, feature_part = line.split(' ', 1)28 labels = list(map(int, label_part.split(',')))29 features = feature_part.strip().split()30 31 for feat in features:32 fid, fval = feat.split(':')33 feat_rows.append(int(fid))34 feat_cols.append(row_idx)35 feat_data.append(float(fval))36 37 for lid in labels:38 label_rows.append(lid)39 label_cols.append(row_idx)40 label_data.append(1.0)41 42 X = csr_matrix((feat_data, (feat_rows, feat_cols)), shape=(num_features, num_samples), dtype=np.float32)43 Y = csr_matrix((label_data, (label_rows, label_cols)), shape=(num_labels, num_samples), dtype=np.uint8)44 return X.transpose(), Y.transpose()45 46# ----------------------- Prediction Functions ----------------------- #47def predict_scores_one_tree(x, tree):48 kmeans1 = tree['level1_kmeans']49 level2_dict = tree['level2_kmeans']50 tree_model = tree['model_tree']51 52 l1_cluster = kmeans1.predict(x.reshape(1, -1))[0]53 54 if l1_cluster not in level2_dict:55 return {}56 57 kmeans2 = level2_dict[l1_cluster]58 l2_cluster = kmeans2.predict(x.reshape(1, -1))[0]59 60 node_key = (l1_cluster, l2_cluster)61 if node_key not in tree_model:62 return {}63 64 classifiers = tree_model[node_key]['classifiers']65 label_ids = tree_model[node_key]['label_ids']66 67 preds = {}68 for clf, lid in zip(classifiers, label_ids):69 score = clf.decision_function(x.reshape(1, -1))[0]70 preds[lid] = score71 return preds72 73def predict_ensemble(X_test_reduced, ensemble):74 all_preds = []75 for i in range(X_test_reduced.shape[0]):76 instance_scores = {}77 x = X_test_reduced[i]78 for tree in ensemble:79 preds = predict_scores_one_tree(x, tree)80 for lid, score in preds.items():81 if lid not in instance_scores:82 instance_scores[lid] = []83 instance_scores[lid].append(score)84 avg_scores = {lid: np.mean(scores) for lid, scores in instance_scores.items()}85 all_preds.append(avg_scores)86 return all_preds87 88def get_topk(pred_scores, k):89 top_labels = []90 for scores in pred_scores:91 if len(scores) == 0:92 top_labels.append([])93 continue94 sorted_labels = sorted(scores.items(), key=lambda x: -x[1])95 top_labels.append([lid for lid, _ in sorted_labels[:k]])96 return top_labels97 98def precision_at_k(preds_topk, Y_true, k):99 hits = 0100 total = len(preds_topk)101 for i, pred_labels in enumerate(preds_topk):102 true_labels = set(Y_true[i].nonzero()[1])103 hits += len(set(pred_labels) & true_labels)104 return hits / (total * k)105 106def evaluate_rforest(ensemble, X_test, Y_test, dim=300):107 svd = TruncatedSVD(n_components=dim, random_state=42)108 X_test_reduced = svd.fit_transform(X_test)109 110 pred_scores = predict_ensemble(X_test_reduced, ensemble)111 112 preds_at_1 = get_topk(pred_scores, k=1)113 preds_at_3 = get_topk(pred_scores, k=3)114 115 p1 = precision_at_k(preds_at_1, Y_test, k=1)116 p3 = precision_at_k(preds_at_3, Y_test, k=3)117 118 y_true = Y_test.toarray()119 y_pred_bin = np.zeros_like(y_true)120 for i, pred in enumerate(get_topk(pred_scores, k=5)):121 y_pred_bin[i, pred] = 1122 123 macro_f1 = f1_score(y_true, y_pred_bin, average='macro', zero_division=0)124 125 return p1, p3, macro_f1126 127# ----------------------- Load Pretrained Model ----------------------- #128@st.cache_resource129def load_model():130 return joblib.load("random_label_forest_model.pkl")131 132ensemble_model = load_model()133 134# ----------------------- User Interface ----------------------- #135uploaded_file = st.file_uploader("Upload Test Dataset (.txt in RCV1 format)", type=["txt"])136 137if uploaded_file:138 with tempfile.NamedTemporaryFile(delete=False, suffix=".txt") as tmp:139 tmp.write(uploaded_file.getvalue())140 test_file_path = tmp.name141 142 if st.button("Test Model"):143 st.info("Parsing dataset and evaluating. Please wait...")144 X_test, Y_test = parse_rcv1_sparse_file(test_file_path)145 p1, p3, macro_f1 = evaluate_rforest(ensemble_model, X_test, Y_test)146 147 st.success("Evaluation Completed!")148 st.metric("Precision@1", f"{p1:.4f}")149 st.metric("Precision@3", f"{p3:.4f}")150 st.metric("Macro F1 Score", f"{macro_f1:.4f}")151 