| 32 | |
| 33 | |
| 34 | class AudioIteratorStreamer(BaseStreamer): |
| 35 | |
| 36 | MAX_TOKEN_LENGTH = 16384 |
| 37 | CHUNK_SIZE = 30 |
| 38 | OVERLAP_SECONDS = 10 |
| 39 | |
| 40 | def __init__( |
| 41 | self, |
| 42 | tokenizer, |
| 43 | model: AsteroidTTSInstruct, |
| 44 | spt: XY_Tokenizer, |
| 45 | device, |
| 46 | use_tqdm: bool = False, |
| 47 | ): |
| 48 | self.tokenizer = tokenizer |
| 49 | self.model = model |
| 50 | self.spt = spt |
| 51 | self.device = device |
| 52 | self.use_tqdm = use_tqdm |
| 53 | self.speech_offset = model.config.speech_token_range[0] |
| 54 | self.channels = model.config.channels |
| 55 | self.duration_code_length = int( |
| 56 | self.CHUNK_SIZE * spt.input_sample_rate // spt.encoder_downsample_rate |
| 57 | ) |
| 58 | self.overlap_code_length = int( |
| 59 | self.OVERLAP_SECONDS * spt.input_sample_rate // spt.encoder_downsample_rate |
| 60 | ) |
| 61 | self.valid_code_length = self.duration_code_length - self.overlap_code_length |
| 62 | self.valid_wav_length = int(self.valid_code_length * spt.decoder_upsample_rate) |
| 63 | print(f"Speech offset: {self.speech_offset}") |
| 64 | print(f"Duration code length: {self.duration_code_length}") |
| 65 | print(f"Overlap code length: {self.overlap_code_length}") |
| 66 | print(f"Valid code length: {self.valid_code_length}") |
| 67 | print(f"Valid wav length: {self.valid_wav_length}") |
| 68 | |
| 69 | self.next_tokens_are_prompt = True |
| 70 | self.token_cache = torch.zeros( |
| 71 | self.MAX_TOKEN_LENGTH, self.channels, dtype=torch.long, device=device |
| 72 | ) |
| 73 | self.token_cache_length = 0 |
| 74 | self.decoded_idx = 0 |
| 75 | self.audio_queue = Queue() |
| 76 | self.stop_signal = None |
| 77 | |
| 78 | # tqdm related |
| 79 | self.pbar = None |
| 80 | if self.use_tqdm: |
| 81 | self.pbar = tqdm(desc="Processing tokens", unit="token", total=None) |
| 82 | |
| 83 | def decode(self, last_chunk=False): |
| 84 | duration_to_decode = min( |
| 85 | self.duration_code_length, |
| 86 | self.token_cache_length - self.decoded_idx - self.channels + 1, |
| 87 | ) |
| 88 | |
| 89 | speech_ids = torch.full((duration_to_decode, self.channels), 0).to(self.device) |
| 90 | for j in range(self.channels): |
| 91 | speech_ids[..., j] = self.token_cache[ |