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()