Spaces:
Sleeping
Sleeping
| """ | |
| FLOAT Streaming Engine - Continuous real-time lipsync generation. | |
| Generates frames continuously in a background thread: | |
| - Idle mode: silent audio → natural idle motion (breathing, blinking, micro-movements) | |
| - Speech mode: TTS audio features injected → lip-synced speech motion | |
| """ | |
| import os | |
| import sys | |
| import math | |
| import time | |
| import threading | |
| import queue | |
| import logging | |
| import random | |
| import numpy as np | |
| import torch | |
| import torch.nn.functional as F | |
| import cv2 | |
| import librosa | |
| logger = logging.getLogger(__name__) | |
| FLOAT_REPO_PATH = "/app/float_repo" | |
| if FLOAT_REPO_PATH not in sys.path: | |
| sys.path.insert(0, FLOAT_REPO_PATH) | |
| class FloatStreamer: | |
| def __init__(self): | |
| self.model = None | |
| self.device = None | |
| self.opt = None | |
| self.s_r = None | |
| self.r_s = None | |
| self.s_r_feats = None | |
| self.prev_x = None | |
| self.prev_wa = None | |
| self.wav2vec_preprocessor = None | |
| self._frame_buffer = queue.Queue(maxsize=75) | |
| self._speech_queue = queue.Queue() | |
| self._running = False | |
| self._thread = None | |
| self._lock = threading.Lock() | |
| self._is_speaking = False | |
| self.fps = 25.0 | |
| self.frames_per_chunk = 50 | |
| self.num_prev_frames = 10 | |
| self.sampling_rate = 16000 | |
| self.ready = False | |
| self.current_dampening = 0.35 | |
| # --- Procedural Idle State Machine Variables --- | |
| self._idle_frame_counter = 0 | |
| self._idle_is_active_state = False | |
| self._idle_state_frames_left = 125 # Start with 5 seconds of calm | |
| # Target values we want to reach | |
| self._idle_target_mid = 0.32 | |
| self._idle_target_amp = 0.06 | |
| # Current values we are smoothly animating | |
| self._idle_current_mid = 0.32 | |
| self._idle_current_amp = 0.06 | |
| def initialize(self, model, device, opt, ref_image_tensor, wav2vec_preprocessor): | |
| self.model = model | |
| self.device = device | |
| self.opt = opt | |
| self.wav2vec_preprocessor = wav2vec_preprocessor | |
| self.fps = opt.fps | |
| self.frames_per_chunk = int(opt.wav2vec_sec * opt.fps) | |
| self.num_prev_frames = opt.num_prev_frames | |
| self.sampling_rate = opt.sampling_rate | |
| self._encode_reference(ref_image_tensor) | |
| with torch.no_grad(): | |
| silence_audio = torch.zeros(1, int(opt.wav2vec_sec * self.sampling_rate)).to(device) | |
| T_silence = self.frames_per_chunk | |
| self._silence_wa = self.model.audio_encoder.inference(silence_audio, seq_len=T_silence) | |
| self._silence_we = self.model.emotion_encoder.predict_emotion(silence_audio).unsqueeze(1) | |
| logger.info(f"[STREAMER] Silence features encoded: wa={self._silence_wa.shape}") | |
| self.prev_x = torch.zeros(1, self.num_prev_frames, opt.dim_w).to(device) | |
| self.prev_wa = torch.zeros(1, self.num_prev_frames, opt.dim_w).to(device) | |
| self.ready = True | |
| logger.info("[STREAMER] Initialized and ready") | |
| def _encode_reference(self, ref_tensor): | |
| with torch.no_grad(): | |
| s = ref_tensor.to(self.device) | |
| self.s_r, r_s_lambda, self.s_r_feats = self.model.encode_image_into_latent(s) | |
| self.r_s = self.model.motion_autoencoder.dec.direction(r_s_lambda) | |
| def update_reference(self, ref_tensor): | |
| with self._lock: | |
| self._encode_reference(ref_tensor) | |
| self.prev_x = torch.zeros(1, self.num_prev_frames, self.opt.dim_w).to(self.device) | |
| self.prev_wa = torch.zeros(1, self.num_prev_frames, self.opt.dim_w).to(self.device) | |
| logger.info("[STREAMER] Reference updated") | |
| def start(self): | |
| if self._running: return | |
| self._running = True | |
| self._thread = threading.Thread(target=self._generation_loop, daemon=True) | |
| self._thread.start() | |
| logger.info("[STREAMER] Generation loop started") | |
| def stop(self): | |
| self._running = False | |
| if self._thread: self._thread.join(timeout=5) | |
| def inject_speech(self, audio_path: str, audio_url: str = None, clear_buffer: bool = True, is_last: bool = True): | |
| t0 = time.time() | |
| if clear_buffer: | |
| self.drain_buffer() | |
| logger.info("[STREAMER] Buffer flushed for immediate speech playback") | |
| speech_array, sr = librosa.load(audio_path, sr=self.sampling_rate) | |
| if is_last: | |
| pad_samples = int(0.5 * self.sampling_rate) | |
| speech_array = np.concatenate([speech_array, np.zeros(pad_samples, dtype=speech_array.dtype)]) | |
| audio_tensor = torch.FloatTensor(speech_array).unsqueeze(0).to(self.device) | |
| with torch.no_grad(): | |
| T = math.ceil(audio_tensor.shape[-1] * self.fps / self.sampling_rate) | |
| wa_full = self.model.audio_encoder.inference(audio_tensor, seq_len=T) | |
| we = self.model.emotion_encoder.predict_emotion(audio_tensor).unsqueeze(1) | |
| num_chunks = math.ceil(T / self.frames_per_chunk) | |
| for i in range(num_chunks): | |
| start_frame = i * self.frames_per_chunk | |
| end_frame = min(start_frame + self.frames_per_chunk, T) | |
| wa_chunk = wa_full[:, start_frame:end_frame] | |
| if wa_chunk.shape[1] < self.frames_per_chunk: | |
| wa_chunk = F.pad(wa_chunk, (0, 0, 0, self.frames_per_chunk - wa_chunk.shape[1]), mode='replicate') | |
| chunk_data = { | |
| "wa": wa_chunk, "we": we, "chunk_index": i, | |
| "total_chunks": num_chunks, "actual_frames": end_frame - start_frame, | |
| } | |
| if i == 0: | |
| chunk_data["speech_start"] = { | |
| "type": "speech_start", "audio_url": audio_url, | |
| "duration": len(speech_array) / self.sampling_rate, "num_chunks": num_chunks, | |
| } | |
| self._speech_queue.put(chunk_data) | |
| if is_last: | |
| self._speech_queue.put({"type": "speech_end"}) | |
| logger.info(f"[STREAMER] Speech injected: {T} frames, {num_chunks} chunks, is_last={is_last}") | |
| def get_frame(self, timeout=0.1): | |
| try: return self._frame_buffer.get(timeout=timeout) | |
| except queue.Empty: return None | |
| def _generation_loop(self): | |
| logger.info("[STREAMER] Generation loop running") | |
| self._idle_cfg_scale = 1.0 | |
| while self._running: | |
| try: | |
| speech_data = None | |
| try: speech_data = self._speech_queue.get_nowait() | |
| except queue.Empty: pass | |
| if speech_data and speech_data.get("type") == "speech_end": | |
| self._is_speaking = False | |
| self._idle_cfg_scale = self.opt.a_cfg_scale | |
| logger.info("[STREAMER] Speech ended, transitioning to idle") | |
| continue | |
| if speech_data and "wa" in speech_data: | |
| self._is_speaking = True | |
| speech_start_event = speech_data.get("speech_start") | |
| self._generate_chunk( | |
| wa=speech_data["wa"], we=speech_data["we"], | |
| actual_frames=speech_data.get("actual_frames", self.frames_per_chunk), | |
| is_speech=True, nfe=5, | |
| speech_start_event=speech_start_event | |
| ) | |
| else: | |
| self._generate_idle_chunk() | |
| except Exception as e: | |
| logger.error(f"[STREAMER] Generation error: {e}", exc_info=True) | |
| time.sleep(0.5) | |
| def _generate_idle_chunk(self): | |
| cfg = self._idle_cfg_scale | |
| if cfg > 1.2: self._idle_cfg_scale = max(1.2, cfg - 0.2) | |
| dampening = torch.zeros(1, self.frames_per_chunk, 1, device=self.device) | |
| for t in range(self.frames_per_chunk): | |
| # 1. State Machine Timer | |
| self._idle_state_frames_left -= 1 | |
| if self._idle_state_frames_left <= 0: | |
| # Toggle state | |
| self._idle_is_active_state = not self._idle_is_active_state | |
| if self._idle_is_active_state: | |
| # Switch to Dynamic/Active (High swaying) | |
| self._idle_target_mid = random.uniform(0.48, 0.52) | |
| self._idle_target_amp = random.uniform(0.12, 0.16) | |
| self._idle_state_frames_left = random.randint(125, 250) # Hold for 5-10s | |
| else: | |
| # Switch to Static/Calm (Subtle breathing) | |
| self._idle_target_mid = random.uniform(0.28, 0.34) | |
| self._idle_target_amp = random.uniform(0.04, 0.08) | |
| self._idle_state_frames_left = random.randint(125, 375) # Hold for 5-15s | |
| # 2. Smooth Interpolation (Lerp) | |
| # This smoothly glides the current values toward the target over ~2 seconds | |
| lerp_speed = 0.02 | |
| self._idle_current_mid += (self._idle_target_mid - self._idle_current_mid) * lerp_speed | |
| self._idle_current_amp += (self._idle_target_amp - self._idle_current_amp) * lerp_speed | |
| # 3. Apply to Sine Wave | |
| global_t = self._idle_frame_counter + t | |
| dampening[0, t, 0] = self._idle_current_mid + self._idle_current_amp * math.sin(global_t * 0.05) | |
| self._idle_frame_counter += self.frames_per_chunk | |
| self._generate_chunk( | |
| wa=self._silence_wa, we=self._silence_we, actual_frames=self.frames_per_chunk, | |
| a_cfg_scale=cfg, e_cfg_scale=1.0, is_speech=False, nfe=3, | |
| dynamic_dampening=dampening | |
| ) | |
| def _generate_chunk(self, wa, we, actual_frames=None, a_cfg_scale=None, e_cfg_scale=None, is_speech=False, nfe=None, speech_start_event=None, dynamic_dampening=None): | |
| if actual_frames is None: actual_frames = self.frames_per_chunk | |
| if a_cfg_scale is None: a_cfg_scale = self.opt.a_cfg_scale | |
| if e_cfg_scale is None: e_cfg_scale = self.opt.e_cfg_scale | |
| if nfe is None: nfe = self.opt.nfe | |
| t0 = time.time() | |
| with self._lock: | |
| r_s = self.r_s | |
| s_r = self.s_r | |
| s_r_feats = self.s_r_feats | |
| prev_x = self.prev_x | |
| prev_wa = self.prev_wa | |
| x0 = torch.randn(1, self.frames_per_chunk, self.opt.dim_w, device=self.device) | |
| time_steps = torch.linspace(0, 1, nfe, device=self.device) | |
| def sample_chunk(tt, zt): | |
| out = self.model.fmt.forward_with_cfv( | |
| t=tt.unsqueeze(0), x=zt, wa=wa, wr=r_s, we=we, | |
| prev_x=prev_x, prev_wa=prev_wa, | |
| a_cfg_scale=a_cfg_scale, r_cfg_scale=self.opt.r_cfg_scale, e_cfg_scale=e_cfg_scale, | |
| ) | |
| return out[:, self.num_prev_frames:] | |
| from torchdiffeq import odeint | |
| trajectory = odeint(sample_chunk, x0, time_steps, atol=self.opt.ode_atol, rtol=self.opt.ode_rtol, method=self.opt.torchdiffeq_ode_method) | |
| sample = trajectory[-1] | |
| t_ode = time.time() - t0 | |
| if not is_speech: | |
| sp = sample.permute(0, 2, 1) | |
| sp = F.avg_pool1d(F.pad(sp, (1, 1), mode='replicate'), kernel_size=3, stride=1) | |
| sample = sp.permute(0, 2, 1) | |
| with self._lock: | |
| self.prev_x = sample[:, -self.num_prev_frames:].clone() | |
| self.prev_wa = wa[:, -self.num_prev_frames:].clone() | |
| if is_speech: | |
| target_dampening = 1.0 | |
| if isinstance(self.current_dampening, torch.Tensor): | |
| start_val = self.current_dampening[0, -1, 0].item() | |
| damp_curve = torch.linspace(start_val, 1.0, sample.shape[1], device=self.device).view(1, -1, 1) | |
| sample_display = sample * damp_curve | |
| elif abs(self.current_dampening - target_dampening) > 0.01: | |
| damp_curve = torch.linspace(self.current_dampening, target_dampening, sample.shape[1], device=self.device).view(1, -1, 1) | |
| sample_display = sample * damp_curve | |
| else: | |
| sample_display = sample * target_dampening | |
| self.current_dampening = 1.0 | |
| else: | |
| if dynamic_dampening is not None: | |
| if isinstance(self.current_dampening, float): | |
| fade = torch.linspace(1.0, 0.0, sample.shape[1], device=self.device).view(1, -1, 1) | |
| damp_curve = self.current_dampening * fade + dynamic_dampening * (1.0 - fade) | |
| sample_display = sample * damp_curve | |
| else: | |
| sample_display = sample * dynamic_dampening | |
| self.current_dampening = dynamic_dampening | |
| else: | |
| sample_display = sample * 0.35 | |
| self.current_dampening = 0.35 | |
| t_dec = time.time() | |
| frames_pushed = 0 | |
| if speech_start_event: | |
| try: | |
| self._frame_buffer.put(speech_start_event, timeout=2.0) | |
| except queue.Full: | |
| logger.warning("[STREAMER] Frame buffer full, dropping speech_start event") | |
| for t in range(min(actual_frames, sample.shape[1])): | |
| if not is_speech and not self._speech_queue.empty(): | |
| if self._frame_buffer.qsize() > 25: | |
| break | |
| s_r_d_t = s_r + sample_display[:, t] | |
| img_t, _ = self.model.motion_autoencoder.dec(s_r_d_t, alpha=None, feats=s_r_feats) | |
| frame = img_t.squeeze().permute(1, 2, 0).detach().clamp(-1, 1) | |
| frame = ((frame + 1) / 2 * 255).to(torch.uint8).cpu().numpy() | |
| _, jpeg_data = cv2.imencode('.jpg', cv2.cvtColor(frame, cv2.COLOR_RGB2BGR), [cv2.IMWRITE_JPEG_QUALITY, 85]) | |
| frame_bytes = jpeg_data.tobytes() | |
| if is_speech: | |
| try: | |
| self._frame_buffer.put(frame_bytes, timeout=2.0) | |
| frames_pushed += 1 | |
| except queue.Full: | |
| pass | |
| else: | |
| try: | |
| self._frame_buffer.put_nowait(frame_bytes) | |
| frames_pushed += 1 | |
| except queue.Full: | |
| pass | |
| t_dec_done = time.time() | |
| def drain_buffer(self): | |
| while not self._frame_buffer.empty(): | |
| try: self._frame_buffer.get_nowait() | |
| except queue.Empty: break | |
| while not self._speech_queue.empty(): | |
| try: self._speech_queue.get_nowait() | |
| except queue.Empty: break | |
| _streamer = None | |
| def get_streamer() -> FloatStreamer: | |
| global _streamer | |
| if _streamer is None: | |
| _streamer = FloatStreamer() | |
| return _streamer |