File size: 1,302 Bytes
c2992d0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""consensus_sampling.py — Consenso via amostragem (Chernoff) (V11.23)."""
from __future__ import annotations
import torch
from collections import Counter

class ConsensusSampling:
    def __init__(self, M=10, temperature=0.8):
        self.M=M; self.temperature=temperature
    def sample(self, model, input_ids, max_new_tokens=30):
        outputs=[]
        model.eval()
        for _ in range(self.M):
            with torch.no_grad():
                gen=model.generate(input_ids, max_new_tokens=max_new_tokens, temperature=self.temperature)
            outputs.append(gen)
        return outputs
    def vote(self, outputs, tokenizer=None):
        if tokenizer:
            texts=[tokenizer.decode(o) for o in outputs]
        else:
            texts=[str(o) for o in outputs]
        counter=Counter(texts)
        consensus, count=counter.most_common(1)[0]
        confidence=count/len(texts)
        return consensus
    def get_consensus(self, model, input_ids, tokenizer=None, max_new_tokens=30):
        outputs=self.sample(model, input_ids, max_new_tokens)
        consensus=self.vote(outputs, tokenizer)
        texts=[tokenizer.decode(o) if tokenizer else str(o) for o in outputs]
        return {"consensus": consensus, "confidence": confidence, "n_samples": self.M, "texts": texts}