{ "cells": [ { "cell_type": "code", "execution_count": 14, "id": "71fa50b4-5593-4a9b-9c2f-7da7c76e82b7", "metadata": {}, "outputs": [ { "name": "stderr", "output_type": "stream", "text": [ "C:\\Users\\Admin\\.cache\\huggingface\\modules\\transformers_modules\\zhihan1996\\DNABERT-2-117M\\7bce263b15377fc15361f52cfab88f8b586abda0\\bert_layers.py:126: UserWarning: Unable to import Triton; defaulting MosaicBERT attention implementation to pytorch (this will reduce throughput when using this model).\n", " warnings.warn(\n" ] }, { "name": "stdout", "output_type": "stream", "text": [ "Saved annotated_result.csv\n" ] } ], "source": [ "import pandas as pd\n", "import re\n", "import torch\n", "from transformers import BertConfig, AutoTokenizer, AutoModelForSequenceClassification\n", "\n", "# 1) Load and clean clustering results\n", "df = pd.read_csv(\"result.csv\")\n", "df['Fragment'] = df['Fragment'].astype(str).apply(lambda x: re.sub(r'[^ACGT]', '', x))\n", "df = df[(df['Fragment'] != '') & df['Fragment'].notna()]\n", "\n", "# Optional: keep original clustering labels (1,2,3)\n", "if 'label' in df.columns:\n", " df['cluster_label_1idx'] = df['label']\n", " df['cluster_label_0idx'] = df['label'] - 1\n", "\n", "# 2) Load model and tokenizer consistently from the same repo\n", "MODEL_NAME = \"Priyasi/MetaTrans_CLM\"\n", "config = BertConfig.from_pretrained(MODEL_NAME, num_labels=3)\n", "tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME, trust_remote_code=True)\n", "model = AutoModelForSequenceClassification.from_pretrained(\n", " MODEL_NAME,\n", " config=config,\n", " trust_remote_code=True,\n", " torch_dtype=\"auto\",\n", ")\n", "model.eval()\n", "\n", "# 3) Predict annotated labels\n", "enc = tokenizer(df['Fragment'].tolist(), truncation=True, padding=True, max_length=512, return_tensors=\"pt\")\n", "with torch.no_grad():\n", " logits = model(**enc).logits\n", " preds = torch.argmax(logits, dim=-1).cpu().numpy()\n", "\n", "df['annotated_label'] = preds # 0-indexed (0,1,2)\n", "df.to_csv(\"annotated_result.csv\", index=False)\n", "print(\"Saved annotated_result.csv\")\n" ] }, { "cell_type": "code", "execution_count": 25, "id": "a11bb5ea-212e-4bf7-9f06-83d51f1bd7f0", "metadata": {}, "outputs": [], "source": [ "from sklearn.model_selection import train_test_split\n", "from transformers import TrainingArguments, Trainer\n", "import numpy as np\n", "from sklearn.metrics import accuracy_score, f1_score, precision_score, recall_score\n", "import torch\n", "\n", "# Use the model’s own num_labels unless you truly need to change it\n", "num_labels = getattr(config, \"num_labels\", None)\n", "if num_labels is None:\n", " num_labels = len(set(df['annotated_label']))\n", " config.num_labels = num_labels # keep config and model aligned\n", "\n", "# Replace labels with annotations\n", "df_train = df.dropna(subset=['Fragment', 'annotated_label']).copy()\n", "df_train['label'] = df_train['annotated_label']\n", "\n", "# Split\n", "train_df, val_df = train_test_split(\n", " df_train, test_size=0.2, random_state=42, stratify=df_train['label']\n", ")\n", "\n", "train_texts = train_df['Fragment'].tolist()\n", "train_labels = train_df['label'].tolist()\n", "val_texts = val_df['Fragment'].tolist()\n", "val_labels = val_df['label'].tolist()\n", "\n", "# Tokenize\n", "train_enc = tokenizer(train_texts, truncation=True, padding=True, max_length=512)\n", "val_enc = tokenizer(val_texts, truncation=True, padding=True, max_length=512)\n", "\n", "class MetaData(torch.utils.data.Dataset):\n", " def __init__(self, encodings, labels):\n", " self.encodings = encodings\n", " self.labels = labels\n", " def __getitem__(self, idx):\n", " item = {k: torch.tensor(v[idx]) for k, v in self.encodings.items()}\n", " item['labels'] = torch.tensor(self.labels[idx])\n", " return item\n", " def __len__(self):\n", " return len(self.labels)\n", "\n", "train_ds = MetaData(train_enc, train_labels)\n", "val_ds = MetaData(val_enc, val_labels)\n", "\n", "# Metrics\n", "def compute_metrics(pred):\n", " labels = pred.label_ids\n", " logits = pred.predictions[0] if isinstance(pred.predictions, tuple) else pred.predictions\n", " preds = np.argmax(logits, axis=-1)\n", " return {\n", " \"accuracy\": accuracy_score(labels, preds),\n", " \"precision\": precision_score(labels, preds, average=\"weighted\"),\n", " \"recall\": recall_score(labels, preds, average=\"weighted\"),\n", " \"f1\": f1_score(labels, preds, average=\"weighted\"),\n", " }\n" ] }, { "cell_type": "code", "execution_count": 31, "id": "7b7225e9-dce1-498b-8207-ce203c6c83c9", "metadata": {}, "outputs": [ { "name": "stderr", "output_type": "stream", "text": [ "C:\\Users\\Admin\\AppData\\Local\\Programs\\Python\\Python38\\lib\\site-packages\\transformers\\training_args.py:1568: FutureWarning: `evaluation_strategy` is deprecated and will be removed in version 4.46 of 🤗 Transformers. Use `eval_strategy` instead\n", " warnings.warn(\n" ] } ], "source": [ "\n", "training_args = TrainingArguments(\n", " output_dir='./results_clm_star',\n", " num_train_epochs=10, \n", " per_device_train_batch_size=16,\n", " per_device_eval_batch_size=32,\n", " warmup_steps=10,\n", " weight_decay=0.01,\n", " learning_rate=2e-5,\n", " logging_dir='./logs_clm_star',\n", " logging_steps=50,\n", " evaluation_strategy=\"epoch\",\n", " save_strategy=\"epoch\",\n", " load_best_model_at_end=True,\n", " metric_for_best_model=\"f1\",\n", " report_to=\"none\",\n", ")\n", "\n", "\n", "trainer = Trainer(\n", " model=model,\n", " args=training_args, # ✅ use the TrainingArguments you defined\n", " train_dataset=train_ds,\n", " eval_dataset=val_ds,\n", " compute_metrics=compute_metrics,\n", ")\n", "\n", "\n" ] }, { "cell_type": "code", "execution_count": 32, "id": "68f04416-7d98-4e23-b013-d9e3f769084b", "metadata": {}, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "--- Starting CLM* fine-tuning on annotated dataset ---\n" ] }, { "data": { "text/html": [ "\n", "
| Epoch | \n", "Training Loss | \n", "Validation Loss | \n", "Accuracy | \n", "Precision | \n", "Recall | \n", "F1 | \n", "
|---|---|---|---|---|---|---|
| 1 | \n", "0.029600 | \n", "0.541489 | \n", "0.926667 | \n", "0.931356 | \n", "0.926667 | \n", "0.926021 | \n", "
| 2 | \n", "0.022900 | \n", "0.693667 | \n", "0.916667 | \n", "0.922330 | \n", "0.916667 | \n", "0.915891 | \n", "
| 3 | \n", "0.025200 | \n", "0.533131 | \n", "0.930000 | \n", "0.931210 | \n", "0.930000 | \n", "0.929878 | \n", "
| 4 | \n", "0.075800 | \n", "0.310411 | \n", "0.953333 | \n", "0.954460 | \n", "0.953333 | \n", "0.953119 | \n", "
| 5 | \n", "0.057900 | \n", "0.324236 | \n", "0.940000 | \n", "0.941670 | \n", "0.940000 | \n", "0.939695 | \n", "
| 6 | \n", "0.024800 | \n", "0.373928 | \n", "0.943333 | \n", "0.946571 | \n", "0.943333 | \n", "0.943028 | \n", "
| 7 | \n", "0.015900 | \n", "0.260695 | \n", "0.963333 | \n", "0.963854 | \n", "0.963333 | \n", "0.963220 | \n", "
| 8 | \n", "0.000800 | \n", "0.425310 | \n", "0.943333 | \n", "0.945108 | \n", "0.943333 | \n", "0.943054 | \n", "
| 9 | \n", "0.000200 | \n", "0.336196 | \n", "0.953333 | \n", "0.954305 | \n", "0.953333 | \n", "0.953182 | \n", "
| 10 | \n", "0.001400 | \n", "0.341137 | \n", "0.950000 | \n", "0.950868 | \n", "0.950000 | \n", "0.949849 | \n", "
"
],
"text/plain": [
"