Spaces:
Paused
Paused
| import numpy as np | |
| import torch | |
| from mediatok.container.gtkv import GtkvReader, VideoTokenBlock, AudioTokenBlock | |
| from mediatok.codecs.base import VideoTokenizer, AudioTokenizer | |
| from mediatok.entropy import entropy_decode | |
| from mediatok.container.gtkv import ENTROPY_CODEC_NAME_MAP | |
| class DecoderPipeline: | |
| def __init__(self, reader: GtkvReader, video_codec: VideoTokenizer, | |
| audio_codec: AudioTokenizer, device: str = "cpu"): | |
| self.reader = reader | |
| self.video_codec = video_codec | |
| self.audio_codec = audio_codec | |
| self.device = torch.device(device) | |
| self.entropy_codec_id = reader.header.entropy_codec_id | |
| self.layer_token_counts = reader.header.layer_token_counts | |
| self.num_layers = reader.header.num_layers | |
| def decode_chunk(self, chunk_index: int, layer_mask: int = 0b111111) -> torch.Tensor: | |
| block = self.reader.read_video_block(chunk_index) | |
| tokens = entropy_decode(block.entropy_payload, self.entropy_codec_id, block.token_count, bits=18) | |
| layer_slices = [] | |
| offset = 0 | |
| for i in range(self.num_layers): | |
| n = block.layer_sizes[i] if i < len(block.layer_sizes) and block.layer_sizes[i] > 0 else self.layer_token_counts[i] | |
| if layer_mask & (1 << i): | |
| layer_slices.append(tokens[offset:offset + n]) | |
| offset += n | |
| frames_per_chunk = self.reader.header.chunk_size_frames or 1 | |
| layer_tensors = [] | |
| for i in range(self.num_layers): | |
| n = block.layer_sizes[i] if i < len(block.layer_sizes) and block.layer_sizes[i] > 0 else 0 | |
| if n > 0 and (layer_mask & (1 << i)): | |
| tpf = n // frames_per_chunk | |
| arr = np.array(layer_slices.pop(0), dtype=np.int64) | |
| layer_tensors.append( | |
| torch.tensor(arr.reshape(1, frames_per_chunk, tpf), dtype=torch.int64, device=self.device) | |
| ) | |
| frames = self.video_codec.decode(layer_tensors) | |
| return frames | |
| def decode_all(self, layer_mask: int = 0b111111) -> list[torch.Tensor]: | |
| frames = [] | |
| for i in range(self.reader.num_chunks): | |
| chunk = self.decode_chunk(i, layer_mask) | |
| frames.append(chunk) | |
| return frames | |
| def decode_audio_chunk(self, chunk_index: int) -> torch.Tensor: | |
| block = self.reader.read_audio_block(chunk_index) | |
| if block is None: | |
| return torch.empty(0, device=self.device) | |
| tokens = entropy_decode(block.entropy_payload, self.entropy_codec_id, | |
| block.codebook * block.frame_count, bits=32) | |
| t = torch.tensor(tokens, dtype=torch.int64, device=self.device) | |
| t = t.view(block.codebook, -1) | |
| audio = self.audio_codec.decode(t) | |
| return audio | |