File size: 14,739 Bytes
a9ae5ea
 
347e66f
 
 
 
a9ae5ea
 
 
 
 
 
 
 
d9d6205
a9ae5ea
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
91317dc
 
 
a9ae5ea
91317dc
 
a9ae5ea
 
 
91317dc
a9ae5ea
 
 
 
 
1a2979f
a9ae5ea
 
91317dc
a9ae5ea
 
 
 
91317dc
d56f673
d9d6205
d56f673
d9d6205
 
 
 
 
 
 
 
 
 
a9ae5ea
 
 
 
 
 
 
 
91317dc
 
 
a9ae5ea
 
 
789cf28
 
 
 
 
 
 
a9ae5ea
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
91317dc
a9ae5ea
 
 
 
 
 
 
91317dc
a9ae5ea
8d85afa
a9ae5ea
4e883e1
 
 
 
 
a9ae5ea
 
8d85afa
 
 
 
a9ae5ea
 
 
 
 
 
 
1ff5b72
a9ae5ea
 
 
 
 
 
 
 
 
1ff5b72
91317dc
 
1ff5b72
 
 
 
91317dc
 
1ff5b72
 
a9ae5ea
8d85afa
 
 
 
a9ae5ea
 
91317dc
 
a9ae5ea
 
 
91317dc
a9ae5ea
 
 
 
91317dc
 
a9ae5ea
 
806cefe
91317dc
a0ae063
a9ae5ea
 
 
806cefe
e1d8608
 
a9ae5ea
91317dc
a9ae5ea
91317dc
9ac0575
a9ae5ea
1a2979f
a9ae5ea
 
 
 
 
 
 
 
a0ae063
347e66f
a0ae063
d56f673
 
 
d9d6205
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d56f673
d9d6205
d56f673
 
 
789cf28
91317dc
d56f673
 
789cf28
a9ae5ea
 
d56f673
91317dc
 
 
 
a9ae5ea
 
 
 
 
 
 
 
 
 
1a2979f
92d6383
a9ae5ea
 
 
91317dc
 
 
a9ae5ea
 
 
 
91317dc
 
a9ae5ea
 
 
347e66f
d9d6205
347e66f
 
 
a9ae5ea
 
 
 
d56f673
 
 
 
 
 
 
 
 
 
 
 
91317dc
d56f673
 
 
 
 
 
 
 
 
 
 
5edcdc7
a9ae5ea
a0ae063
9ac0575
 
 
 
 
 
 
a9ae5ea
91317dc
a0ae063
91317dc
 
a0ae063
91317dc
a9ae5ea
 
 
 
91317dc
a0ae063
 
 
 
 
 
 
8d85afa
a0ae063
 
 
 
 
91317dc
a9ae5ea
 
 
 
 
91317dc
 
a9ae5ea
91317dc
 
a9ae5ea
 
 
 
 
 
 
94cdccd
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
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
"""
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
        )

    @torch.no_grad()
    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