File size: 1,756 Bytes
fac1a3a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
# CUDA_VISIBLE_DEVICES=0,1,2,3 python classify_generics.py
import pandas as pd
from tqdm import tqdm
from transformers import RobertaForSequenceClassification, RobertaTokenizer

import random
random.seed(101)

## PARAMETERS ##

input_csv = "data/input.csv"
text_column = "text_column_name"
output_csv = ""
output_csv = f"{input_csv}_scored_for_generics.csv" if not output_csv else output_csv


model_path = "ilyocoris/generics-classifier-mgen"


## LOAD DATA & MODEL ##

data = pd.read_csv(input_csv)
data = data.to_dict(orient="records")

def load_roberta_classifier(model_name, checkpoint=None, base_model="roberta-large"):
    # model_path = f"data/training_runs/{model_name}/models/checkpoint-{checkpoint}" if checkpoint else f"data/models/{model_name}"
    model_path = f"{model_name}/checkpoint-{checkpoint}" if checkpoint else model_name
    model = RobertaForSequenceClassification.from_pretrained(
        model_path,
        device_map="auto",
        num_labels=1
    )
    tokenizer = RobertaTokenizer.from_pretrained(base_model)
    return model, tokenizer

def classify_batch(batch_sentences, model, tokenizer):
    inputs = tokenizer(batch_sentences, return_tensors="pt", padding=True, truncation=True)
    inputs = {k: v.cuda() for k, v in inputs.items()}
    outputs = model(**inputs)
    return outputs.logits.squeeze().cpu().detach().numpy()

model, tokenizer = load_roberta_classifier(model_path) 

## RUN CLASSIFIER ##

scores = []
batch_size = 16
for i in tqdm(range(0, len(data), batch_size)):
    batch = data[i:i+batch_size]
    sentences = [d[text_column]  for d in batch]
    scores.extend(classify_batch(sentences, model, tokenizer).tolist())

df = pd.DataFrame(data)
df["score"] = scores
df.to_csv(f"{output_csv}", index=False)