| 97 | |
| 98 | |
| 99 | class Model( |
| 100 | nn.Module, |
| 101 | PyTorchModelHubMixin, |
| 102 | repo_url="https://github.com/SesameAILabs/csm", |
| 103 | pipeline_tag="text-to-speech", |
| 104 | license="apache-2.0", |
| 105 | ): |
| 106 | def __init__(self, config: ModelArgs): |
| 107 | super().__init__() |
| 108 | self.config = config |
| 109 | |
| 110 | self.backbone, backbone_dim = _prepare_transformer(FLAVORS[config.backbone_flavor]()) |
| 111 | self.decoder, decoder_dim = _prepare_transformer(FLAVORS[config.decoder_flavor]()) |
| 112 | |
| 113 | self.text_embeddings = nn.Embedding(config.text_vocab_size, backbone_dim) |
| 114 | self.audio_embeddings = nn.Embedding(config.audio_vocab_size * config.audio_num_codebooks, backbone_dim) |
| 115 | |
| 116 | self.projection = nn.Linear(backbone_dim, decoder_dim, bias=False) |
| 117 | self.codebook0_head = nn.Linear(backbone_dim, config.audio_vocab_size, bias=False) |
| 118 | self.audio_head = nn.Parameter(torch.empty(config.audio_num_codebooks - 1, decoder_dim, config.audio_vocab_size)) |
| 119 | |
| 120 | def setup_caches(self, max_batch_size: int) -> torch.Tensor: |
| 121 | """Setup KV caches and return a causal mask.""" |
| 122 | dtype = next(self.parameters()).dtype |
| 123 | device = next(self.parameters()).device |
| 124 | |
| 125 | with device: |
| 126 | self.backbone.setup_caches(max_batch_size, dtype) |
| 127 | self.decoder.setup_caches(max_batch_size, dtype, decoder_max_seq_len=self.config.audio_num_codebooks) |
| 128 | |
| 129 | self.register_buffer("backbone_causal_mask", _create_causal_mask(self.backbone.max_seq_len, device)) |
| 130 | self.register_buffer("decoder_causal_mask", _create_causal_mask(self.config.audio_num_codebooks, device)) |
| 131 | |
| 132 | def generate_frame( |
| 133 | self, |
| 134 | tokens: torch.Tensor, |
| 135 | tokens_mask: torch.Tensor, |
| 136 | input_pos: torch.Tensor, |
| 137 | temperature: float, |
| 138 | topk: int, |
| 139 | ) -> torch.Tensor: |
| 140 | """ |
| 141 | Args: |
| 142 | tokens: (batch_size, seq_len, audio_num_codebooks+1) |
| 143 | tokens_mask: (batch_size, seq_len, audio_num_codebooks+1) |
| 144 | input_pos: (batch_size, seq_len) positions for each token |
| 145 | mask: (batch_size, seq_len, max_seq_len |
| 146 | |
| 147 | Returns: |
| 148 | (batch_size, audio_num_codebooks) sampled tokens |
| 149 | """ |
| 150 | dtype = next(self.parameters()).dtype |
| 151 | b, s, _ = tokens.size() |
| 152 | |
| 153 | assert self.backbone.caches_are_enabled(), "backbone caches are not enabled" |
| 154 | curr_backbone_mask = _index_causal_mask(self.backbone_causal_mask, input_pos) |
| 155 | embeds = self._embed_tokens(tokens) |
| 156 | masked_embeds = embeds * tokens_mask.unsqueeze(-1) |
nothing calls this directly
no outgoing calls
no test coverage detected