MCPcopy Create free account
hub / github.com/SesameAILabs/csm / Model

Class Model

models.py:99–203  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

97
98
99class 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)

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected