import struct from pathlib import Path import numpy as np import torch import soundfile as sf from mediatok.container.gtkv import GtkvWriter, GtkvHeader, VideoTokenBlock, AudioTokenBlock from mediatok.codecs.base import VideoTokenizer, AudioTokenizer from mediatok.entropy import entropy_encode from mediatok.container.gtkv import VIDEO_TOKENIZER_ID_MAP, AUDIO_TOKENIZER_ID_MAP, ENTROPY_CODEC_ID_MAP def _choose_entropy_codec() -> int: return 4 class EncoderPipeline: def __init__(self, video_codec: VideoTokenizer, audio_codec: AudioTokenizer, chunk_size_frames: int = 128, entropy_codec: str = "rans"): self.video_codec = video_codec self.audio_codec = audio_codec self.chunk_size_frames = chunk_size_frames self.entropy_codec_id = ENTROPY_CODEC_ID_MAP.get(entropy_codec, 2) def encode_file(self, video_path: str, audio_path: str, output_path: str, width: int, height: int, fps: int = 30): header = GtkvHeader( video_tokenizer_id=VIDEO_TOKENIZER_ID_MAP.get(self.video_codec.name.split("-")[0], 0), audio_tokenizer_id=AUDIO_TOKENIZER_ID_MAP.get(self.audio_codec.name.split("-")[0], 0), width=width, height=height, fps=fps, audio_sample_rate=self.audio_codec.sample_rate, chunk_size_frames=self.chunk_size_frames, entropy_codec_id=self.entropy_codec_id, num_layers=self.video_codec.num_layers, layer_token_counts=self.video_codec.layer_token_counts[:6], ) with GtkvWriter(output_path, header) as writer: self._encode_video_chunks(writer, video_path, header) self._encode_audio_chunks(writer, audio_path, header) def _encode_video_chunks(self, writer: GtkvWriter, video_path: str, header: GtkvHeader): import numpy as np import torch n_frames = header.num_video_frames if header.num_video_frames else 0 if n_frames == 0: return for chunk_start in range(0, n_frames, self.chunk_size_frames): chunk_end = min(chunk_start + self.chunk_size_frames, n_frames) n = chunk_end - chunk_start dummy_frames = torch.randn(1, 3, n, header.height, header.width) tokens = self.video_codec.encode(dummy_frames) flat_tokens = [] layer_sizes = [0] * 6 for i, t in enumerate(tokens): t_np = t.cpu().numpy().ravel().astype(np.int32).tolist() flat_tokens.extend(t_np) if i < 6: layer_sizes[i] = len(t_np) entropy_payload = entropy_encode(flat_tokens, self.entropy_codec_id, bits=18) block = VideoTokenBlock( token_count=len(flat_tokens), layer_sizes=layer_sizes, tokens=flat_tokens, entropy_payload=entropy_payload, ) writer.write_chunk(block) def _encode_audio_chunks(self, writer: GtkvWriter, audio_path: str, header: GtkvHeader): import torch import soundfile as sf if not audio_path or not Path(audio_path).exists(): return data, sr = sf.read(audio_path) if sr != self.audio_codec.sample_rate: import numpy as np from scipy import signal ratio = self.audio_codec.sample_rate / sr new_len = int(len(data) * ratio) data = signal.resample(data, new_len) audio_t = torch.from_numpy(data).float() if audio_t.dim() == 1: audio_t = audio_t.unsqueeze(0) tokens = self.audio_codec.encode(audio_t) t_np = tokens.cpu().numpy().ravel().astype(np.int32).tolist() payload = entropy_encode(t_np, self.entropy_codec_id, bits=32) block = AudioTokenBlock( codebook=self.audio_codec.num_codebooks, frame_count=header.num_video_frames, tokens=t_np, entropy_payload=payload, ) writer.write_chunk(block)