Text Classification
Transformers
Safetensors
bert
multilingual
multi-label-classification
community-notes
topic-classification
text-embeddings-inference
Instructions to use ychuai/community-notes-topic-classifier with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use ychuai/community-notes-topic-classifier with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-classification", model="ychuai/community-notes-topic-classifier")# Load model directly from transformers import AutoTokenizer, AutoModelForSequenceClassification tokenizer = AutoTokenizer.from_pretrained("ychuai/community-notes-topic-classifier") model = AutoModelForSequenceClassification.from_pretrained("ychuai/community-notes-topic-classifier", device_map="auto") - Notebooks
- Google Colab
- Kaggle
| """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') | |
| 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)) | |