File size: 3,207 Bytes
724479b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
from __future__ import annotations

import torch
from torch import nn
from transformers.models.cohere2.modeling_cohere2 import Cohere2ForCausalLM

from .configuration_tinyaya import TinyAyaConfig


class TinyAyaForCausalLM(Cohere2ForCausalLM):
    config_class = TinyAyaConfig

    def __init__(self, config):
        super().__init__(config)
        hidden = int(config.hidden_size)
        self.stop_predictor = nn.Sequential(
            nn.LayerNorm(hidden),
            nn.Linear(hidden, max(64, hidden // 4)),
            nn.GELU(),
            nn.Linear(max(64, hidden // 4), 1),
        )

    @staticmethod
    def _sample(scores, do_sample, temperature, top_k):
        if not do_sample:
            return scores.argmax(dim=-1, keepdim=True)
        scores = scores / max(float(temperature), 1e-5)
        if int(top_k) > 0:
            k = min(int(top_k), scores.shape[-1])
            cutoff = torch.topk(scores, k, dim=-1).values[:, -1:]
            scores = scores.masked_fill(scores < cutoff, torch.finfo(scores.dtype).min)
        return torch.multinomial(torch.softmax(scores.float(), dim=-1), 1)

    @torch.inference_mode()
    def generate_audio(
        self,
        input_ids,
        attention_mask,
        allowed_ids,
        max_new_tokens=2048,
        min_new_tokens=8,
        temperature=0.8,
        top_k=30,
        do_sample=True,
    ):
        out = self(
            input_ids=input_ids,
            attention_mask=attention_mask,
            use_cache=True,
            output_hidden_states=True,
            return_dict=True,
        )
        emitted = input_ids
        cache = out.past_key_values
        scores = out.logits[:, -1, :]
        hidden = out.hidden_states[-1][:, -1:, :]
        mask = attention_mask
        allowed_ids = allowed_ids.to(scores.device)
        audio_end_id = int(self.config.audio_end_token_id)
        for step in range(int(max_new_tokens)):
            if step >= int(min_new_tokens):
                stop = torch.sigmoid(self.stop_predictor(hidden).squeeze(-1))
                if bool((stop > 0.5).all()):
                    eos = input_ids.new_full((input_ids.shape[0], 1), audio_end_id)
                    return torch.cat((emitted, eos), dim=1)
            restricted = torch.full_like(scores, torch.finfo(scores.dtype).min)
            restricted.index_copy_(1, allowed_ids, scores.index_select(1, allowed_ids))
            if step < int(min_new_tokens):
                restricted[:, audio_end_id] = torch.finfo(scores.dtype).min
            token = self._sample(restricted, do_sample, temperature, top_k)
            emitted = torch.cat((emitted, token), dim=1)
            if bool((token == audio_end_id).all()):
                break
            mask = torch.cat((mask, torch.ones_like(token)), dim=1)
            out = self(
                input_ids=token,
                attention_mask=mask,
                past_key_values=cache,
                use_cache=True,
                output_hidden_states=True,
                return_dict=True,
            )
            cache = out.past_key_values
            scores = out.logits[:, -1, :]
            hidden = out.hidden_states[-1][:, -1:, :]
        return emitted