Spaces:
Paused
Paused
File size: 2,811 Bytes
92076a7 | 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 | 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
|