Skip to main content

rivet/multigpu/
single_file.rs

1//! Single-file multi-GPU orchestration: [`run_multigpu_single_file`] + chunk
2//! worker spawn helper.
3
4use std::sync::Arc;
5use std::sync::atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering};
6
7use anyhow::{Result, anyhow, bail};
8use tokio::sync::{Notify, mpsc};
9use tokio::task::JoinSet;
10
11use codec::encode::EncodedPacket;
12use codec::frame::VideoCodec;
13
14use crate::encoder_worker::{
15    ChunkPackets, EncoderWorkerConfig, RungCodecInvariant, run_chunk_encoder_worker_blocking,
16};
17use crate::frame_queue::SegmentChunkQueue;
18use crate::gpu_pool::GpuLease;
19use crate::progress::ProgressSink;
20use crate::spec::Rung;
21
22use crate::cmaf_util::total_segments_for_rung;
23
24use super::{
25    FANOUT_CHANNEL_CAPACITY, HELPER_POLL_INTERVAL, QUEUE_CAPACITY, MultiGpuParams, WorkerCtx,
26    report, spawn_progress_reporter,
27};
28
29/// One rung's full ordered AV1 packet stream, stitched from chunks encoded
30/// across GPUs. The caller muxes these into a single MP4 (+ audio).
31#[derive(Debug)]
32pub struct RungPackets {
33    pub rung_index: usize,
34    pub codec: VideoCodec,
35    pub width: u32,
36    pub height: u32,
37    pub label: String,
38    pub packets: Vec<EncodedPacket>,
39}
40
41/// Single-file counterpart to [`run_multigpu_hls`]: decode once, fan to per-rung
42/// scalers, and dynamically schedule each rung's GOP-sized chunks across all
43/// GPUs (fair lease pool + mid-flight helper dispatch + cross-vendor codec
44/// invariant). Each worker encodes its chunk to packets (a fresh encoder per
45/// chunk → first frame is an IDR); the finalizer concatenates them in segment
46/// order into one ordered packet stream per rung — no disk round-trip.
47pub async fn run_multigpu_single_file(
48    params: MultiGpuParams<'_>,
49    sink: Arc<dyn ProgressSink>,
50) -> Result<Vec<Option<RungPackets>>> {
51    let rungs = params.rungs;
52    let n = rungs.len();
53    if n == 0 {
54        return Ok(Vec::new());
55    }
56    let total_segments = total_segments_for_rung(params.total_input_frames, params.keyframe_interval);
57    if total_segments == 0 {
58        bail!(
59            "multigpu single-file: total_segments == 0 (frames={}, keyframe_interval={})",
60            params.total_input_frames,
61            params.keyframe_interval
62        );
63    }
64
65    // Pre-flight encoder probe (same fail-fast as the HLS path).
66    {
67        let probe = codec::encode::EncoderConfig {
68            width: rungs[0].width,
69            height: rungs[0].height,
70            frame_rate: params.frame_rate,
71            gpu_index: None,
72            codec: params.codec,
73            ..Default::default()
74        };
75        codec::encode::select_encoder(probe, None).map_err(|e| {
76            anyhow!(
77                "no {:?} encoder available on this host ({e}); need NVENC / AMF / QSV, or build \
78                 with the `ffmpeg` feature",
79                params.codec
80            )
81        })?;
82    }
83
84    tracing::info!(
85        rungs = n,
86        total_segments,
87        gpu_pool_capacity = params.gpu_pool.capacity(),
88        "multi-GPU single-file phase starting"
89    );
90
91    let queues: Vec<Arc<SegmentChunkQueue>> =
92        (0..n).map(|_| Arc::new(SegmentChunkQueue::new(QUEUE_CAPACITY))).collect();
93    let frames_encoded: Vec<Arc<AtomicU64>> = (0..n).map(|_| Arc::new(AtomicU64::new(0))).collect();
94    let scaler_active: Vec<Arc<AtomicBool>> =
95        (0..n).map(|_| Arc::new(AtomicBool::new(false))).collect();
96    let rung_invariants: Vec<Arc<std::sync::RwLock<Option<RungCodecInvariant>>>> =
97        (0..n).map(|_| Arc::new(std::sync::RwLock::new(None))).collect();
98    // Per-rung packet collectors (each its own Arc so chunk workers can push).
99    let contributions: Vec<Arc<std::sync::Mutex<Vec<ChunkPackets>>>> =
100        (0..n).map(|_| Arc::new(std::sync::Mutex::new(Vec::new()))).collect();
101    let active_workers: Arc<Vec<AtomicUsize>> =
102        Arc::new((0..n).map(|_| AtomicUsize::new(0)).collect());
103    let rung_done: Arc<Vec<Notify>> = Arc::new((0..n).map(|_| Notify::new()).collect());
104    let finalized: Arc<Vec<AtomicBool>> = Arc::new((0..n).map(|_| AtomicBool::new(false)).collect());
105
106    let progress_stop = Arc::new(AtomicBool::new(false));
107    let progress_handle = spawn_progress_reporter(
108        rungs.to_vec(),
109        frames_encoded.clone(),
110        finalized.clone(),
111        params.total_input_frames,
112        Arc::clone(&sink),
113        Arc::clone(&progress_stop),
114    );
115
116    // Finalizers: stitch each rung's chunks (sorted, deduped) into one stream.
117    let total_input_frames = params.total_input_frames;
118    let codec = params.codec; // Copy; captured by each finalizer closure
119    let (finalizer_tx, mut finalizer_rx) =
120        mpsc::channel::<(usize, Result<Option<RungPackets>>)>(n.max(1));
121    let mut finalizer_handles = Vec::with_capacity(n);
122    for idx in 0..n {
123        let collector = Arc::clone(&contributions[idx]);
124        let active_h = Arc::clone(&active_workers);
125        let rung_done_h = Arc::clone(&rung_done);
126        let finalized_h = Arc::clone(&finalized);
127        let tx = finalizer_tx.clone();
128        let rung = rungs[idx].clone();
129        let total_segments = total_segments;
130        let sink = Arc::clone(&sink);
131        finalizer_handles.push(tokio::spawn(async move {
132            loop {
133                let notified = rung_done_h[idx].notified();
134                if active_h[idx].load(Ordering::Acquire) == 0 {
135                    break;
136                }
137                notified.await;
138            }
139            let mut chunks: Vec<ChunkPackets> = std::mem::take(&mut *collector.lock().unwrap());
140            if chunks.is_empty() {
141                finalized_h[idx].store(true, Ordering::Release);
142                let _ = tx.send((idx, Ok(None))).await;
143                return;
144            }
145            chunks.sort_by_key(|c| c.segment_idx);
146            chunks.dedup_by_key(|c| c.segment_idx);
147            // Coverage: contiguous 0..total_segments.
148            let got = chunks.len();
149            let contiguous = chunks
150                .iter()
151                .enumerate()
152                .all(|(i, c)| c.segment_idx == i);
153            let result = if got != total_segments as usize || !contiguous {
154                Err(anyhow!(
155                    "rung {} chunk coverage incomplete: expected {} contiguous chunks, got {}",
156                    rung.label,
157                    total_segments,
158                    got
159                ))
160            } else {
161                let mut packets: Vec<EncodedPacket> = Vec::new();
162                for c in chunks {
163                    packets.extend(c.packets);
164                }
165                let bytes: u64 = packets.iter().map(|p| p.data.len() as u64).sum();
166                report(
167                    sink.as_ref(),
168                    idx,
169                    &rung,
170                    crate::progress::RungStatus::Completed,
171                    total_input_frames,
172                    Some(total_input_frames),
173                    got as u32,
174                    bytes,
175                    None,
176                );
177                Ok(Some(RungPackets {
178                    rung_index: idx,
179                    codec,
180                    width: rung.width,
181                    height: rung.height,
182                    label: rung.label.clone(),
183                    packets,
184                }))
185            };
186            finalized_h[idx].store(true, Ordering::Release);
187            let _ = tx.send((idx, result)).await;
188        }));
189    }
190    drop(finalizer_tx);
191
192    let mut indexed: Vec<(usize, Rung)> = rungs.iter().cloned().enumerate().collect();
193    indexed.sort_by_key(|(_, r)| r.short_side());
194
195    // Decode pump(s) + fan-out.
196    let mut frame_senders = Vec::with_capacity(n);
197    let mut frame_receivers: Vec<Option<tokio::sync::mpsc::Receiver<codec::frame::VideoFrame>>> =
198        Vec::with_capacity(n);
199    for _ in 0..n {
200        let (tx, rx) = tokio::sync::mpsc::channel(FANOUT_CHANNEL_CAPACITY);
201        frame_senders.push(tx);
202        frame_receivers.push(Some(rx));
203    }
204    let use_shared_pump = n <= params.gpu_pool.capacity();
205    let mut pump_tasks: JoinSet<Result<u64>> = JoinSet::new();
206    if use_shared_pump {
207        let clips = params.clip_sources_for(params.decode_gpu_for(0));
208        let senders = frame_senders;
209        let rt = tokio::runtime::Handle::current();
210        pump_tasks.spawn(async move {
211            tokio::task::spawn_blocking(move || {
212                crate::decode_pump::run_spliced_decode_pump_blocking(clips, senders, rt)
213            })
214            .await
215            .map_err(|e| anyhow!("shared pump join error: {e}"))
216            .and_then(|r| r)
217        });
218    } else {
219        for (idx, sender) in frame_senders.into_iter().enumerate() {
220            let clips = params.clip_sources_for(params.decode_gpu_for(idx));
221            let rt = tokio::runtime::Handle::current();
222            pump_tasks.spawn(async move {
223                tokio::task::spawn_blocking(move || {
224                    crate::decode_pump::run_spliced_decode_pump_blocking(clips, vec![sender], rt)
225                })
226                .await
227                .map_err(|e| anyhow!("per-rung pump {idx} join error: {e}"))
228                .and_then(|r| r)
229            });
230        }
231    }
232
233    // Per-rung scalers.
234    let mut scaler_tasks: JoinSet<(usize, Result<usize>)> = JoinSet::new();
235    for (idx, rung) in rungs.iter().cloned().enumerate() {
236        let rx = frame_receivers[idx].take().expect("scaler rx slot");
237        let cfg = crate::rung_scaler::RungScalerConfig {
238            rung_idx: idx,
239            target_width: rung.width,
240            target_height: rung.height,
241            frames_per_chunk: params.keyframe_interval,
242        };
243        let queue = Arc::clone(&queues[idx]);
244        let rt = tokio::runtime::Handle::current();
245        let scaler_flag = Arc::clone(&scaler_active[idx]);
246        let active_h = Arc::clone(&active_workers);
247        let rung_done_h = Arc::clone(&rung_done);
248        scaler_flag.store(true, Ordering::Release);
249        active_h[idx].fetch_add(1, Ordering::AcqRel);
250        scaler_tasks.spawn(async move {
251            let result = tokio::task::spawn_blocking(move || {
252                crate::rung_scaler::run_rung_scaler_blocking(cfg, rx, queue, rt)
253            })
254            .await
255            .map_err(|e| anyhow!("scaler join error: {e}"))
256            .and_then(|r| r);
257            scaler_flag.store(false, Ordering::Release);
258            let prev = active_h[idx].fetch_sub(1, Ordering::AcqRel);
259            if prev == 1 {
260                rung_done_h[idx].notify_one();
261            }
262            (idx, result)
263        });
264    }
265
266    // Initial chunk workers.
267    let mut worker_tasks: JoinSet<(usize, Result<()>)> = JoinSet::new();
268    let ctx = WorkerCtx {
269        codec: params.codec,
270        frame_rate: params.frame_rate,
271        output_color_metadata: params.output_color_metadata,
272        output_pixel_format: params.output_pixel_format,
273        timescale: params.timescale,
274        per_frame_ticks: params.per_frame_ticks,
275        keyframe_interval: params.keyframe_interval,
276        segment_target_ticks: params.segment_target_ticks,
277        output_root: params.output_root.clone(),
278        constant_qp: params.constant_qp,
279    };
280    for (idx, rung) in indexed.iter().cloned() {
281        let lease = match Arc::clone(&params.gpu_pool).claim().await {
282            Some(l) => l,
283            None => {
284                progress_stop.store(true, Ordering::Release);
285                let _ = progress_handle.await;
286                bail!("multigpu single-file: GPU pool returned no lease; at least one GPU required");
287            }
288        };
289        spawn_chunk_worker(
290            &ctx,
291            idx,
292            &rung,
293            Arc::clone(&queues[idx]),
294            Arc::clone(&frames_encoded[idx]),
295            lease,
296            Arc::clone(&contributions[idx]),
297            Arc::clone(&active_workers),
298            Arc::clone(&rung_done),
299            Arc::clone(&rung_invariants[idx]),
300            Some(&mut worker_tasks),
301        );
302    }
303
304    // Helper dispatcher.
305    let helper_cancel = Arc::new(AtomicBool::new(false));
306    let helper_handle = {
307        let cancel = Arc::clone(&helper_cancel);
308        let pool = Arc::clone(&params.gpu_pool);
309        let queues = queues.clone();
310        let scaler_active = scaler_active.clone();
311        let frames_encoded = frames_encoded.clone();
312        let contributions = contributions.clone();
313        let active_workers = Arc::clone(&active_workers);
314        let rung_done = Arc::clone(&rung_done);
315        let rung_invariants = rung_invariants.clone();
316        let rungs_owned: Vec<Rung> = rungs.to_vec();
317        let ctx = ctx.clone();
318        tokio::spawn(async move {
319            loop {
320                if cancel.load(Ordering::Acquire) {
321                    break;
322                }
323                tokio::time::sleep(HELPER_POLL_INTERVAL).await;
324                if pool.pending_claimers() > 0 {
325                    continue;
326                }
327                let mut target = None;
328                for (idx, q) in queues.iter().enumerate() {
329                    let scaler_alive = scaler_active[idx].load(Ordering::Acquire);
330                    let has_pending = q.pushed_segments() > q.popped_segments();
331                    if scaler_alive || has_pending {
332                        target = Some(idx);
333                        break;
334                    }
335                }
336                let Some(rung_idx) = target else { break };
337                let lease = match pool.try_claim() {
338                    Some(l) => l,
339                    None => continue,
340                };
341                tracing::info!(rung_idx, gpu_index = lease.gpu_index, "single-file helper dispatch");
342                spawn_chunk_worker(
343                    &ctx,
344                    rung_idx,
345                    &rungs_owned[rung_idx],
346                    Arc::clone(&queues[rung_idx]),
347                    Arc::clone(&frames_encoded[rung_idx]),
348                    lease,
349                    Arc::clone(&contributions[rung_idx]),
350                    Arc::clone(&active_workers),
351                    Arc::clone(&rung_done),
352                    Arc::clone(&rung_invariants[rung_idx]),
353                    None,
354                );
355            }
356        })
357    };
358
359    // Drain.
360    let mut completed: Vec<Option<RungPackets>> = (0..n).map(|_| None).collect();
361    let mut pumps_remaining = pump_tasks.len();
362    let mut scalers_remaining = n;
363    let mut workers_remaining = n;
364    let mut finalizers_remaining = n;
365    macro_rules! teardown_err {
366        ($e:expr) => {{
367            helper_cancel.store(true, Ordering::Release);
368            let _ = helper_handle.await;
369            progress_stop.store(true, Ordering::Release);
370            let _ = progress_handle.await;
371            return Err($e);
372        }};
373    }
374    while pumps_remaining > 0 || scalers_remaining > 0 || workers_remaining > 0 || finalizers_remaining > 0 {
375        tokio::select! {
376            biased;
377            p = pump_tasks.join_next(), if pumps_remaining > 0 => match p {
378                Some(Ok(Ok(_))) => pumps_remaining -= 1,
379                Some(Ok(Err(e))) => teardown_err!(anyhow!("decode pump failed: {e}")),
380                Some(Err(je)) => teardown_err!(anyhow!("pump join error: {je}")),
381                None => pumps_remaining = 0,
382            },
383            s = scaler_tasks.join_next(), if scalers_remaining > 0 => match s {
384                Some(Ok((_, Ok(_)))) => scalers_remaining -= 1,
385                Some(Ok((idx, Err(e)))) => teardown_err!(anyhow!("scaler {idx} failed: {e}")),
386                Some(Err(je)) => teardown_err!(anyhow!("scaler join error: {je}")),
387                None => scalers_remaining = 0,
388            },
389            w = worker_tasks.join_next(), if workers_remaining > 0 => match w {
390                Some(Ok((_, Ok(())))) => workers_remaining -= 1,
391                Some(Ok((idx, Err(e)))) => teardown_err!(anyhow!("chunk worker for rung {idx} failed: {e}")),
392                Some(Err(je)) => teardown_err!(anyhow!("worker join error: {je}")),
393                None => workers_remaining = 0,
394            },
395            f = finalizer_rx.recv(), if finalizers_remaining > 0 => match f {
396                Some((idx, Ok(opt))) => { completed[idx] = opt; finalizers_remaining -= 1; }
397                Some((idx, Err(e))) => teardown_err!(anyhow!("finalizer for rung {idx} failed: {e}")),
398                None => finalizers_remaining = 0,
399            },
400        }
401    }
402    helper_cancel.store(true, Ordering::Release);
403    let _ = helper_handle.await;
404    progress_stop.store(true, Ordering::Release);
405    let _ = progress_handle.await;
406    for h in finalizer_handles {
407        let _ = h.await;
408    }
409    Ok(completed)
410}
411
412#[allow(clippy::too_many_arguments)]
413fn spawn_chunk_worker(
414    ctx: &WorkerCtx,
415    rung_idx: usize,
416    rung: &Rung,
417    queue: Arc<SegmentChunkQueue>,
418    frames_encoded: Arc<AtomicU64>,
419    lease: GpuLease,
420    collector: Arc<std::sync::Mutex<Vec<ChunkPackets>>>,
421    active_workers: Arc<Vec<AtomicUsize>>,
422    rung_done: Arc<Vec<Notify>>,
423    rung_invariant: Arc<std::sync::RwLock<Option<RungCodecInvariant>>>,
424    worker_tasks: Option<&mut JoinSet<(usize, Result<()>)>>,
425) {
426    let gpu_index = lease.gpu_index;
427    let gpu_vendor = lease.vendor;
428    let cfg = EncoderWorkerConfig {
429        rung_idx,
430        codec: ctx.codec,
431        width: rung.width,
432        height: rung.height,
433        frame_rate: ctx.frame_rate,
434        quality: rung.quality.crf.unwrap_or(codec::encode::AUTO_FROM_TARGET),
435        speed_preset: rung.quality.speed_preset.unwrap_or(codec::encode::AUTO_FROM_TARGET),
436        target: rung.quality.target,
437        tier: rung.quality.tier,
438        threads: 0,
439        gpu_index: Some(gpu_index),
440        gpu_vendor: Some(gpu_vendor),
441        output_color_metadata: ctx.output_color_metadata,
442        output_pixel_format: ctx.output_pixel_format,
443        constant_qp: ctx.constant_qp,
444        timescale: ctx.timescale,
445        per_frame_ticks: ctx.per_frame_ticks,
446        keyframe_interval: ctx.keyframe_interval,
447        segment_target_ticks: ctx.segment_target_ticks,
448        output_dir: ctx.output_root.clone(),
449        rung_invariant,
450    };
451    active_workers[rung_idx].fetch_add(1, Ordering::AcqRel);
452    let body = async move {
453        let (progress_tx, mut progress_rx) = mpsc::channel::<u64>(32);
454        let cfg_for_worker = cfg.clone();
455        let queue_for_worker = Arc::clone(&queue);
456        let rt = tokio::runtime::Handle::current();
457        let counter = Arc::clone(&frames_encoded);
458        let out = Arc::clone(&collector);
459        let blocking = tokio::task::spawn_blocking(move || {
460            run_chunk_encoder_worker_blocking(cfg_for_worker, queue_for_worker, rt, counter, progress_tx, out)
461        });
462        let drain = async move { while progress_rx.recv().await.is_some() {} };
463        let (_, br) = tokio::join!(drain, blocking);
464        let task_status: Result<()> = match br {
465            Ok(Ok(())) => Ok(()),
466            Ok(Err(e)) => Err(e),
467            Err(e) => Err(anyhow!("chunk worker join error: {e}")),
468        };
469        drop(lease);
470        let prev = active_workers[rung_idx].fetch_sub(1, Ordering::AcqRel);
471        if prev == 1 {
472            rung_done[rung_idx].notify_one();
473        }
474        (rung_idx, task_status)
475    };
476    match worker_tasks {
477        Some(set) => {
478            set.spawn(body);
479        }
480        None => {
481            tokio::spawn(async move {
482                let _ = body.await;
483            });
484        }
485    }
486}