CoolFace
Apppublic

doby4u/chattts

sourceHugging Facemitupdated 2y agoView on Hugging Face
2likes
infer_utils.py45 linesDownload Raw Back to utils
1 2import torch3import torch.nn.functional as F4 5    6class CustomRepetitionPenaltyLogitsProcessorRepeat():7 8    def __init__(self, penalty: float, max_input_ids, past_window):9        if not isinstance(penalty, float) or not (penalty > 0):10            raise ValueError(f"`penalty` has to be a strictly positive float, but is {penalty}")11 12        self.penalty = penalty13        self.max_input_ids = max_input_ids14        self.past_window = past_window15 16    def __call__(self, input_ids: torch.LongTensor, scores: torch.FloatTensor) -> torch.FloatTensor:17        18        input_ids = input_ids[:, -self.past_window:]19        freq = F.one_hot(input_ids, scores.size(1)).sum(1)20        freq[self.max_input_ids:] = 021        alpha = self.penalty**freq22        scores = torch.where(scores < 0, scores*alpha, scores/alpha)23 24        return scores25    26class CustomRepetitionPenaltyLogitsProcessor():27 28    def __init__(self, penalty: float, max_input_ids, past_window):29        if not isinstance(penalty, float) or not (penalty > 0):30            raise ValueError(f"`penalty` has to be a strictly positive float, but is {penalty}")31 32        self.penalty = penalty33        self.max_input_ids = max_input_ids34        self.past_window = past_window35 36    def __call__(self, input_ids: torch.LongTensor, scores: torch.FloatTensor) -> torch.FloatTensor:37        38        input_ids = input_ids[:, -self.past_window:]39        score = torch.gather(scores, 1, input_ids)40        _score = score.detach().clone()41        score = torch.where(score < 0, score * self.penalty, score / self.penalty)42        score[input_ids>=self.max_input_ids] = _score[input_ids>=self.max_input_ids]43        scores.scatter_(1, input_ids, score)44        45        return scores