viclickbait_gnn / src /extract_embeddings.py
minhy112's picture
Upload viclickbait_gnn project
877049d verified
Raw
History Blame Contribute Delete
5.44 kB
import argparse
import json
from pathlib import Path
import numpy as np
import pandas as pd
import torch
from sklearn.decomposition import TruncatedSVD
from sklearn.feature_extraction.text import TfidfVectorizer
from torch.utils.data import DataLoader
from transformers import AutoModel, AutoTokenizer
from datasets import TitleDataset
from model_phobert import PhoBERTClassifier
from utils import ensure_dir, resolve_device, set_seed
def extract_transformer_embeddings(
dataframe: pd.DataFrame,
model_name: str,
max_length: int,
batch_size: int,
pooling: str,
device: torch.device,
checkpoint_path: str | None = None,
) -> np.ndarray:
tokenizer = AutoTokenizer.from_pretrained(model_name, use_fast=False)
if checkpoint_path:
model = PhoBERTClassifier(model_name=model_name, pooling=pooling).to(device)
state_dict = torch.load(checkpoint_path, map_location="cpu")
model.load_state_dict(state_dict)
model.eval()
encoder = None
else:
model = AutoModel.from_pretrained(model_name).to(device)
model.eval()
encoder = model
dataset = TitleDataset(dataframe, tokenizer=tokenizer, max_length=max_length)
loader = DataLoader(dataset, batch_size=batch_size, shuffle=False)
vectors: list[np.ndarray] = []
with torch.no_grad():
for batch in loader:
input_ids = batch["input_ids"].to(device)
attention_mask = batch["attention_mask"].to(device)
if checkpoint_path:
pooled = model.encode(input_ids=input_ids, attention_mask=attention_mask)
else:
outputs = encoder(input_ids=input_ids, attention_mask=attention_mask)
hidden = outputs.last_hidden_state
if pooling == "mean":
mask = attention_mask.unsqueeze(-1).expand(hidden.size()).float()
pooled = (hidden * mask).sum(dim=1) / torch.clamp(mask.sum(dim=1), min=1e-9)
else:
pooled = hidden[:, 0, :]
vectors.append(pooled.cpu().numpy().astype(np.float32))
return np.concatenate(vectors, axis=0)
def extract_tfidf_svd_embeddings(dataframe: pd.DataFrame, max_features: int, n_components: int) -> np.ndarray:
vectorizer = TfidfVectorizer(max_features=max_features, ngram_range=(1, 2), sublinear_tf=True)
matrix = vectorizer.fit_transform(dataframe["title"].tolist())
svd = TruncatedSVD(n_components=n_components, random_state=42)
vectors = svd.fit_transform(matrix)
return vectors.astype(np.float32)
def main() -> None:
parser = argparse.ArgumentParser(description="Extract node features for graph models.")
parser.add_argument("--input", required=True)
parser.add_argument("--output", required=True, help="Path to .npy file.")
parser.add_argument("--metadata", default=None)
parser.add_argument("--model", default="vinai/phobert-base")
parser.add_argument("--backend", choices=["transformer", "tfidf_svd"], default="transformer")
parser.add_argument("--pooling", choices=["cls", "mean"], default="cls")
parser.add_argument("--checkpoint", default=None, help="Optional fine-tuned PhoBERT classifier checkpoint.")
parser.add_argument("--batch_size", type=int, default=16)
parser.add_argument("--max_length", type=int, default=64)
parser.add_argument("--svd_components", type=int, default=256)
parser.add_argument("--tfidf_max_features", type=int, default=12000)
parser.add_argument("--device", default=None)
parser.add_argument("--seed", type=int, default=42)
args = parser.parse_args()
set_seed(args.seed)
dataframe = pd.read_csv(args.input).sort_values("node_id").reset_index(drop=True)
device = resolve_device(args.device)
if args.backend == "transformer":
features = extract_transformer_embeddings(
dataframe=dataframe,
model_name=args.model,
max_length=args.max_length,
batch_size=args.batch_size,
pooling=args.pooling,
device=device,
checkpoint_path=args.checkpoint,
)
backend_metadata = {
"backend": "transformer",
"model": args.model,
"pooling": args.pooling,
"checkpoint": args.checkpoint,
}
else:
features = extract_tfidf_svd_embeddings(
dataframe=dataframe,
max_features=args.tfidf_max_features,
n_components=args.svd_components,
)
backend_metadata = {
"backend": "tfidf_svd",
"tfidf_max_features": args.tfidf_max_features,
"svd_components": args.svd_components,
}
output_path = Path(args.output)
ensure_dir(output_path.parent)
np.save(output_path, features)
metadata = {
"num_nodes": int(features.shape[0]),
"feature_dim": int(features.shape[1]),
"node_ids": dataframe["node_id"].tolist(),
"ids": dataframe["id"].tolist() if "id" in dataframe.columns else None,
**backend_metadata,
}
metadata_path = Path(args.metadata) if args.metadata else output_path.with_suffix(".json")
with metadata_path.open("w", encoding="utf-8") as file:
json.dump(metadata, file, indent=2, ensure_ascii=False)
print(f"Saved node features to {output_path} with shape {features.shape}")
if __name__ == "__main__":
main()