BiGRU_T_version / src /bigru_t /reasoning /consensus_sampling.py
PowerMachine's picture
V6.5: 1000 samples (500+500), MTP val+entropy, EWC+W8A8 eval investigation, integrate gru-ring modules
c2992d0 verified
Raw History Blame Contribute Delete
1.3 kB
"""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}