CoolFace
Apppublic

muffin2006/document-classification-env

sourceHugging Faceupdated 6mo agoView on Hugging Face
1likes
baseline_inference.py74 linesDownload Raw Back to root
1import argparse2import json3import time4import pickle5import os6from sklearn.feature_extraction.text import TfidfVectorizer7from sklearn.linear_model import LogisticRegression8from sklearn.pipeline import Pipeline9from tasks import TaskDataGenerator10from environment import DocumentClassificationEnv11 12MODEL_DIR = os.path.dirname(os.path.abspath(__file__))13 14def get_model_path(difficulty):15    return os.path.join(MODEL_DIR, f"model_{difficulty}.pkl")16 17def train_model(difficulty):18    print(f"Training {difficulty} model...")19    all_texts, all_labels = [], []20    for seed in range(10):21        gen = TaskDataGenerator(difficulty, seed=seed)22        docs, labels = gen.generate_task_data()23        for doc, label in zip(docs, labels):24            all_texts.append(doc["content"])25            all_labels.append(label)26    pipeline = Pipeline([27        ("tfidf", TfidfVectorizer(ngram_range=(1,2), max_features=5000, sublinear_tf=True)),28        ("clf", LogisticRegression(max_iter=1000, C=5.0, solver="lbfgs"))29    ])30    pipeline.fit(all_texts, all_labels)31    with open(get_model_path(difficulty), "wb") as f:32        pickle.dump(pipeline, f)33    print(f"  Saved!")34    return pipeline35 36def load_or_train(difficulty):37    path = get_model_path(difficulty)38    if os.path.exists(path):39        with open(path, "rb") as f:40            return pickle.load(f)41    return train_model(difficulty)42 43def run_task(difficulty):44    model = load_or_train(difficulty)45    env = DocumentClassificationEnv(task_difficulty=difficulty, seed=42)46    obs, _ = env.reset()47    correct, total, t0, terminated = 0, 0, time.time(), False48    while not terminated:49        action = int(model.predict([obs["content"]])[0])50        obs, reward, terminated, _, info = env.step(action)51        correct += int(info.get("is_correct", False))52        total += 153    return correct/total if total > 0 else 0, time.time()-t054 55def main():56    parser = argparse.ArgumentParser()57    parser.add_argument("--task", default="all")58    parser.add_argument("--output", default="baseline_results.json")59    args = parser.parse_args()60    tasks = ["easy","medium","hard"] if args.task=="all" else [args.task]61    results = {}62    for task in tasks:63        print(f"\nRunning {task.upper()} task...")64        score, elapsed = run_task(task)65        results[task] = {"score": round(score,4), "time": round(elapsed,2)}66        print(f"  {task.upper()} Score: {score:.4f} ({elapsed:.1f}s)")67    print("\n=== BASELINE SCORES ===")68    for t,r in results.items():69        print(f"  {t.upper()} Score: {r['score']:.4f}")70    with open(args.output,"w") as f:71        json.dump(results,f,indent=2)72 73if __name__ == "__main__":74    main()