File size: 2,776 Bytes
cec807d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
"""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))