ychuai's picture
upload
cec807d verified
Raw
History Blame Contribute Delete
2.78 kB
"""Predict post-level topics with the training pipeline's overlapping-window pooling."""
import argparse
import json
from pathlib import Path
import torch
from transformers import AutoTokenizer, AutoModelForSequenceClassification
def combine_text(post_text='', summaries=()):
post_text = post_text.strip()
summaries = list(dict.fromkeys(s.strip() for s in summaries if s.strip()))
sections = ['Original post:\n'+post_text] if post_text else []
sections += [f'Community Note {i}:\n{s}' for i,s in enumerate(summaries,1)]
if not sections:
raise ValueError('Provide post text or at least one nonempty note summary')
return '\n\n'.join(sections)
class TopicPredictor:
def __init__(self, model_path, device='cpu'):
path=Path(model_path)
self.settings=json.loads((path/'topic_config.json').read_text())
self.tokenizer=AutoTokenizer.from_pretrained(path,local_files_only=True,use_fast=True)
self.model=AutoModelForSequenceClassification.from_pretrained(path,local_files_only=True).to(device).eval()
self.device=device
if self.settings['pooling']!='max_logits':
raise ValueError('Unsupported pooling configuration')
@torch.inference_mode()
def predict(self, post_text='', summaries=(), window_batch_size=8):
if window_batch_size < 1:
raise ValueError('window_batch_size must be positive')
text=combine_text(post_text,summaries)
encoded=self.tokenizer(text,truncation=True,max_length=self.settings['max_length'],
stride=self.settings['stride'],return_overflowing_tokens=True)
windows=[{k:encoded[k][i] for k in self.tokenizer.model_input_names if k in encoded}
for i in range(len(encoded['input_ids']))]
pooled=None
for start in range(0,len(windows),window_batch_size):
batch=self.tokenizer.pad(windows[start:start+window_batch_size],return_tensors='pt').to(self.device)
logits=self.model(**batch).logits.max(dim=0).values
pooled=logits if pooled is None else torch.maximum(pooled,logits)
scores=pooled.sigmoid().cpu().tolist()
return [{'topic':topic,'score':score,'selected':score>=self.settings['threshold']}
for topic,score in zip(self.settings['categories'],scores)]
if __name__=='__main__':
parser=argparse.ArgumentParser(description=__doc__)
parser.add_argument('--model',default=str(Path(__file__).resolve().parent))
parser.add_argument('--post',default='')
parser.add_argument('--note',action='append',default=[])
parser.add_argument('--device',default='cpu')
args=parser.parse_args()
print(json.dumps(TopicPredictor(args.model,args.device).predict(args.post,args.note),indent=2))