mediatok-player / pipeline /decoder.py
Daankular's picture
Upload folder using huggingface_hub
92076a7 verified
Raw
History Blame Contribute Delete
2.81 kB
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