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