BabakScrapes's picture
Upload pipeline.py
8cc90b3 verified
Raw
History Blame Contribute Delete
8.59 kB
from __future__ import annotations
import re
from pathlib import Path
from typing import List, Tuple
import matplotlib.pyplot as plt
import numpy as np
import torch
from transformers import (
AutoModelForSequenceClassification,
AutoModelForTokenClassification,
AutoTokenizer,
)
BASE_DIR = Path(__file__).resolve().parent
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
CLAUSE_MODEL_DIR = BASE_DIR / "clause_model_512"
CLASSIFICATION_MODEL_DIR = BASE_DIR / "classfication_model"
tokenizer = AutoTokenizer.from_pretrained(
"roberta-base",
use_fast=True,
add_prefix_space=True,
)
clause_model = AutoModelForTokenClassification.from_pretrained(
str(CLAUSE_MODEL_DIR)
).to(DEVICE).eval()
classification_model = AutoModelForSequenceClassification.from_pretrained(
str(CLASSIFICATION_MODEL_DIR)
).to(DEVICE).eval()
labels2attrs = {
"##BOUNDED EVENT (SPECIFIC)": ("specific", "dynamic", "episodic"),
"##BOUNDED EVENT (GENERIC)": ("generic", "dynamic", "episodic"),
"##UNBOUNDED EVENT (SPECIFIC)": ("specific", "dynamic", "static"),
"##UNBOUNDED EVENT (GENERIC)": ("generic", "dynamic", "static"),
"##BASIC STATE": ("specific", "stative", "static"),
"##COERCED STATE (SPECIFIC)": ("specific", "dynamic", "static"),
"##COERCED STATE (GENERIC)": ("generic", "dynamic", "static"),
"##PERFECT COERCED STATE (SPECIFIC)": ("specific", "dynamic", "episodic"),
"##PERFECT COERCED STATE (GENERIC)": ("generic", "dynamic", "episodic"),
"##GENERIC SENTENCE (DYNAMIC)": ("generic", "dynamic", "habitual"),
"##GENERIC SENTENCE (STATIC)": ("generic", "stative", "static"),
"##GENERIC SENTENCE (HABITUAL)": ("generic", "stative", "habitual"),
"##GENERALIZING SENTENCE (DYNAMIC)": ("specific", "dynamic", "habitual"),
"##GENERALIZING SENTENCE (STATIVE)": ("specific", "stative", "habitual"),
"##QUESTION": ("NA", "NA", "NA"),
"##IMPERATIVE": ("NA", "NA", "NA"),
"##NONSENSE": ("NA", "NA", "NA"),
"##OTHER": ("NA", "NA", "NA"),
}
label_names = list(labels2attrs.keys())
index2label = {i: label for i, label in enumerate(label_names)}
def split_sentences(text: str) -> List[str]:
text = re.sub(r"\s+", " ", text).strip()
if not text:
return []
# Lightweight sentence splitter to avoid spaCy dependency.
sentences = re.split(r"(?<=[.!?])\s+", text)
return [s.strip() for s in sentences if s.strip()]
def auto_split(text: str, max_words: int = 200) -> List[str]:
sentences = split_sentences(text)
if not sentences:
return []
snippets: List[str] = []
current_words: List[str] = []
for sentence in sentences:
sent_words = sentence.split()
if current_words and len(current_words) + len(sent_words) > max_words:
snippets.append(" ".join(current_words).strip())
current_words = sent_words[:]
else:
current_words.extend(sent_words)
if current_words:
snippets.append(" ".join(current_words).strip())
return snippets
def majority_vote(values: List[int]) -> int:
if not values:
return 1
counts = np.bincount(values)
return int(np.argmax(counts))
@torch.inference_mode()
def get_pred_clause_labels(text: str) -> List[int]:
words = text.strip().split()
if not words:
return []
enc = tokenizer(
words,
is_split_into_words=True,
return_tensors="pt",
truncation=True,
max_length=512,
padding="max_length",
)
word_ids = enc.word_ids(batch_index=0)
model_inputs = {k: v.to(DEVICE) for k, v in enc.items()}
logits = clause_model(**model_inputs).logits[0]
token_preds = logits.argmax(dim=-1).detach().cpu().tolist()
aligned_preds: List[List[int]] = [[] for _ in words]
for token_idx, word_id in enumerate(word_ids):
if word_id is None:
continue
aligned_preds[word_id].append(token_preds[token_idx])
return [majority_vote(preds) if preds else 1 for preds in aligned_preds]
def seg_clause(text: str) -> List[str]:
words = text.strip().split()
if not words:
return []
labels = get_pred_clause_labels(text)
segmented_clauses: List[List[str]] = []
prev_label = 2
current_clause: List[str] | None = None
for word, label in zip(words, labels):
if prev_label == 2:
current_clause = []
if current_clause is not None:
current_clause.append(word)
if label == 2 and prev_label in [0, 1]:
segmented_clauses.append(current_clause[:])
current_clause = None
prev_label = label
if current_clause:
segmented_clauses.append(current_clause[:])
return [" ".join(clause) for clause in segmented_clauses if clause]
@torch.inference_mode()
def get_pred_classification_labels(
clauses: List[str], batch_size: int = 32
) -> List[Tuple[str, Tuple[str, str, str]]]:
results: List[Tuple[str, Tuple[str, str, str]]] = []
for i in range(0, len(clauses), batch_size):
batch = clauses[i : i + batch_size]
enc = tokenizer(
batch,
return_tensors="pt",
truncation=True,
max_length=128,
padding="max_length",
)
model_inputs = {k: v.to(DEVICE) for k, v in enc.items()}
logits = classification_model(**model_inputs).logits
pred_ids = logits.argmax(dim=-1).detach().cpu().tolist()
pred_labels = [index2label[idx] for idx in pred_ids]
results.extend((clause, labels2attrs[label]) for clause, label in zip(batch, pred_labels))
return results
def label_visualization(
clause2labels: List[Tuple[str, Tuple[str, str, str]]]
):
total_clauses = len(clause2labels)
if total_clauses == 0:
fig = plt.figure(figsize=(10, 4))
plt.text(0.5, 0.5, "No clauses detected.", ha="center", va="center")
plt.axis("off")
return fig
aspect_labels, genericity_labels, boundedness_labels = [], [], []
for _, attrs in clause2labels:
genericity_label, aspect_label, boundedness_label = attrs
genericity_labels.append(genericity_label)
aspect_labels.append(aspect_label)
boundedness_labels.append(boundedness_label)
aspect_dict = {
"Dynamic": aspect_labels.count("dynamic"),
"Stative": aspect_labels.count("stative"),
"NA": aspect_labels.count("NA"),
}
genericity_dict = {
"Generic": genericity_labels.count("generic"),
"Specific": genericity_labels.count("specific"),
"NA": genericity_labels.count("NA"),
}
boundedness_dict = {
"Static": boundedness_labels.count("static"),
"Episodic": boundedness_labels.count("episodic"),
"Habitual": boundedness_labels.count("habitual"),
"NA": boundedness_labels.count("NA"),
}
def proportions(d: dict[str, int]) -> tuple[list[str], list[float]]:
filtered = {k: v / total_clauses for k, v in d.items() if v > 0}
return list(filtered.keys()), list(filtered.values())
fig, axs = plt.subplots(1, 3, figsize=(10, 6))
fig.tight_layout(pad=5.0)
labels, values = proportions(aspect_dict)
axs[0].pie(values, labels=labels, autopct="%.0f%%", normalize=True)
axs[0].set_title("Eventivity")
labels, values = proportions(genericity_dict)
axs[1].pie(values, labels=labels, autopct="%.0f%%", normalize=True)
axs[1].set_title("Genericity")
labels, values = proportions(boundedness_dict)
axs[2].pie(values, labels=labels, autopct="%.0f%%", normalize=True)
axs[2].set_title("Boundedness/Habituality")
return fig
def run_pipeline(text: str):
text = (text or "").strip()
if not text:
empty_fig = label_visualization([])
return [], [], empty_fig
snippets = auto_split(text)
all_clauses: List[str] = []
for snippet in snippets:
all_clauses.extend(seg_clause(snippet))
clause2labels = get_pred_classification_labels(all_clauses)
output_clauses = [(clause, str(i + 1)) for i, clause in enumerate(all_clauses)]
highlighted_attrs = [(clause, str(attrs)) for clause, attrs in clause2labels]
figure = label_visualization(clause2labels)
return output_clauses, highlighted_attrs, figure