Spaces:
Paused
Paused
| 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) | |