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