Skip to main content

subspace_farmer/plotter/
gpu.rs

1//! GPU plotter
2
3mod gpu_encoders_manager;
4pub mod metrics;
5pub mod wgpu;
6
7use crate::plotter::gpu::gpu_encoders_manager::GpuRecordsEncoderManager;
8use crate::plotter::gpu::metrics::GpuPlotterMetrics;
9use crate::plotter::{Plotter, SectorPlottingProgress};
10use async_lock::{Mutex as AsyncMutex, Semaphore, SemaphoreGuardArc};
11use async_trait::async_trait;
12use bytes::Bytes;
13use event_listener_primitives::{Bag, HandlerId};
14use futures::channel::mpsc;
15use futures::stream::FuturesUnordered;
16use futures::{FutureExt, Sink, SinkExt, StreamExt, select, stream};
17use prometheus_client::registry::Registry;
18use std::error::Error;
19use std::fmt;
20use std::future::pending;
21use std::num::TryFromIntError;
22use std::pin::pin;
23use std::sync::Arc;
24use std::sync::atomic::{AtomicBool, Ordering};
25use std::task::Poll;
26use std::time::Instant;
27use subspace_core_primitives::PublicKey;
28use subspace_core_primitives::sectors::SectorIndex;
29use subspace_data_retrieval::piece_getter::PieceGetter;
30use subspace_erasure_coding::ErasureCoding;
31use subspace_farmer_components::FarmerProtocolInfo;
32use subspace_farmer_components::plotting::{
33    DownloadSectorOptions, EncodeSectorOptions, PlottingError, RecordsEncoder, download_sector,
34    encode_sector, write_sector,
35};
36use subspace_kzg::Kzg;
37use subspace_process::AsyncJoinOnDrop;
38use tokio::task::yield_now;
39use tracing::{Instrument, warn};
40
41/// Type alias used for event handlers
42pub type HandlerFn3<A, B, C> = Arc<dyn Fn(&A, &B, &C) + Send + Sync + 'static>;
43type Handler3<A, B, C> = Bag<HandlerFn3<A, B, C>, A, B, C>;
44
45#[derive(Default, Debug)]
46struct Handlers {
47    plotting_progress: Handler3<PublicKey, SectorIndex, SectorPlottingProgress>,
48}
49
50/// GPU-specific [`RecordsEncoder`] with extra APIs
51pub trait GpuRecordsEncoder: RecordsEncoder + fmt::Debug + Send {
52    /// GPU encoder type, typically related to GPU vendor
53    const TYPE: &'static str;
54}
55
56/// GPU plotter
57pub struct GpuPlotter<PG, GRE> {
58    piece_getter: PG,
59    downloading_semaphore: Arc<Semaphore>,
60    gpu_records_encoders_manager: GpuRecordsEncoderManager<GRE>,
61    global_mutex: Arc<AsyncMutex<()>>,
62    kzg: Kzg,
63    erasure_coding: ErasureCoding,
64    handlers: Arc<Handlers>,
65    tasks_sender: mpsc::Sender<AsyncJoinOnDrop<()>>,
66    _background_tasks: AsyncJoinOnDrop<()>,
67    abort_early: Arc<AtomicBool>,
68    metrics: Option<Arc<GpuPlotterMetrics>>,
69}
70
71impl<PG, GRE> fmt::Debug for GpuPlotter<PG, GRE>
72where
73    GRE: GpuRecordsEncoder + 'static,
74{
75    #[inline]
76    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
77        f.debug_struct(&format!("GpuPlotter[type = {}]", GRE::TYPE))
78            .finish_non_exhaustive()
79    }
80}
81
82impl<PG, RE> Drop for GpuPlotter<PG, RE> {
83    #[inline]
84    fn drop(&mut self) {
85        self.abort_early.store(true, Ordering::Release);
86        self.tasks_sender.close_channel();
87    }
88}
89
90#[async_trait]
91impl<PG, GRE> Plotter for GpuPlotter<PG, GRE>
92where
93    PG: PieceGetter + Clone + Send + Sync + 'static,
94    GRE: GpuRecordsEncoder + 'static,
95{
96    async fn has_free_capacity(&self) -> Result<bool, String> {
97        Ok(self.downloading_semaphore.try_acquire().is_some())
98    }
99
100    async fn plot_sector(
101        &self,
102        public_key: PublicKey,
103        sector_index: SectorIndex,
104        farmer_protocol_info: FarmerProtocolInfo,
105        pieces_in_sector: u16,
106        _replotting: bool,
107        progress_sender: mpsc::Sender<SectorPlottingProgress>,
108    ) {
109        let start = Instant::now();
110
111        // Done outside the future below as a backpressure, ensuring that it is not possible to
112        // schedule unbounded number of plotting tasks
113        let downloading_permit = self.downloading_semaphore.acquire_arc().await;
114
115        self.plot_sector_internal(
116            start,
117            downloading_permit,
118            public_key,
119            sector_index,
120            farmer_protocol_info,
121            pieces_in_sector,
122            progress_sender,
123        )
124        .await
125    }
126
127    async fn try_plot_sector(
128        &self,
129        public_key: PublicKey,
130        sector_index: SectorIndex,
131        farmer_protocol_info: FarmerProtocolInfo,
132        pieces_in_sector: u16,
133        _replotting: bool,
134        progress_sender: mpsc::Sender<SectorPlottingProgress>,
135    ) -> bool {
136        let start = Instant::now();
137
138        let Some(downloading_permit) = self.downloading_semaphore.try_acquire_arc() else {
139            return false;
140        };
141
142        self.plot_sector_internal(
143            start,
144            downloading_permit,
145            public_key,
146            sector_index,
147            farmer_protocol_info,
148            pieces_in_sector,
149            progress_sender,
150        )
151        .await;
152
153        true
154    }
155}
156
157impl<PG, GRE> GpuPlotter<PG, GRE>
158where
159    PG: PieceGetter + Clone + Send + Sync + 'static,
160    GRE: GpuRecordsEncoder + 'static,
161{
162    /// Create new instance.
163    ///
164    /// Returns an error if empty list of encoders is provided.
165    pub fn new(
166        piece_getter: PG,
167        downloading_semaphore: Arc<Semaphore>,
168        gpu_records_encoders: Vec<GRE>,
169        global_mutex: Arc<AsyncMutex<()>>,
170        kzg: Kzg,
171        erasure_coding: ErasureCoding,
172        registry: Option<&mut Registry>,
173    ) -> Result<Self, TryFromIntError> {
174        let (tasks_sender, mut tasks_receiver) = mpsc::channel(1);
175
176        // Basically runs plotting tasks in the background and allows to abort on drop
177        let background_tasks = AsyncJoinOnDrop::new(
178            tokio::spawn(async move {
179                let background_tasks = FuturesUnordered::new();
180                let mut background_tasks = pin!(background_tasks);
181                // Just so that `FuturesUnordered` will never end
182                background_tasks.push(AsyncJoinOnDrop::new(tokio::spawn(pending::<()>()), true));
183
184                loop {
185                    select! {
186                        maybe_background_task = tasks_receiver.next().fuse() => {
187                            let Some(background_task) = maybe_background_task else {
188                                break;
189                            };
190
191                            background_tasks.push(background_task);
192                        },
193                        _ = background_tasks.select_next_some() => {
194                            // Nothing to do
195                        }
196                    }
197                }
198            }),
199            true,
200        );
201
202        let abort_early = Arc::new(AtomicBool::new(false));
203        let gpu_records_encoders_manager = GpuRecordsEncoderManager::new(gpu_records_encoders)?;
204        let metrics = registry.map(|registry| {
205            Arc::new(GpuPlotterMetrics::new(
206                registry,
207                GRE::TYPE,
208                gpu_records_encoders_manager.gpu_records_encoders(),
209            ))
210        });
211
212        Ok(Self {
213            piece_getter,
214            downloading_semaphore,
215            gpu_records_encoders_manager,
216            global_mutex,
217            kzg,
218            erasure_coding,
219            handlers: Arc::default(),
220            tasks_sender,
221            _background_tasks: background_tasks,
222            abort_early,
223            metrics,
224        })
225    }
226
227    /// Subscribe to plotting progress notifications
228    pub fn on_plotting_progress(
229        &self,
230        callback: HandlerFn3<PublicKey, SectorIndex, SectorPlottingProgress>,
231    ) -> HandlerId {
232        self.handlers.plotting_progress.add(callback)
233    }
234
235    #[allow(clippy::too_many_arguments)]
236    async fn plot_sector_internal<PS>(
237        &self,
238        start: Instant,
239        downloading_permit: SemaphoreGuardArc,
240        public_key: PublicKey,
241        sector_index: SectorIndex,
242        farmer_protocol_info: FarmerProtocolInfo,
243        pieces_in_sector: u16,
244        mut progress_sender: PS,
245    ) where
246        PS: Sink<SectorPlottingProgress> + Unpin + Send + 'static,
247        PS::Error: Error,
248    {
249        if let Some(metrics) = &self.metrics {
250            metrics.sector_plotting.inc();
251        }
252
253        let progress_updater = ProgressUpdater {
254            public_key,
255            sector_index,
256            handlers: Arc::clone(&self.handlers),
257            metrics: self.metrics.clone(),
258        };
259
260        let plotting_fut = {
261            let piece_getter = self.piece_getter.clone();
262            let gpu_records_encoders_manager = self.gpu_records_encoders_manager.clone();
263            let global_mutex = Arc::clone(&self.global_mutex);
264            let kzg = self.kzg.clone();
265            let erasure_coding = self.erasure_coding.clone();
266            let abort_early = Arc::clone(&self.abort_early);
267            let metrics = self.metrics.clone();
268
269            async move {
270                // Downloading
271                let downloaded_sector = {
272                    if !progress_updater
273                        .update_progress_and_events(
274                            &mut progress_sender,
275                            SectorPlottingProgress::Downloading,
276                        )
277                        .await
278                    {
279                        return;
280                    }
281
282                    // Take mutex briefly to make sure plotting is allowed right now
283                    global_mutex.lock().await;
284
285                    let downloading_start = Instant::now();
286
287                    let downloaded_sector_fut = download_sector(DownloadSectorOptions {
288                        public_key: &public_key,
289                        sector_index,
290                        piece_getter: &piece_getter,
291                        farmer_protocol_info,
292                        kzg: &kzg,
293                        erasure_coding: &erasure_coding,
294                        pieces_in_sector,
295                    });
296
297                    let downloaded_sector = match downloaded_sector_fut.await {
298                        Ok(downloaded_sector) => downloaded_sector,
299                        Err(error) => {
300                            warn!(%error, "Failed to download sector");
301
302                            progress_updater
303                                .update_progress_and_events(
304                                    &mut progress_sender,
305                                    SectorPlottingProgress::Error {
306                                        error: format!("Failed to download sector: {error}"),
307                                    },
308                                )
309                                .await;
310
311                            return;
312                        }
313                    };
314
315                    if !progress_updater
316                        .update_progress_and_events(
317                            &mut progress_sender,
318                            SectorPlottingProgress::Downloaded(downloading_start.elapsed()),
319                        )
320                        .await
321                    {
322                        return;
323                    }
324
325                    downloaded_sector
326                };
327
328                // Plotting
329                let (sector, plotted_sector) = {
330                    let mut records_encoder = gpu_records_encoders_manager.get_encoder().await;
331                    if let Some(metrics) = &metrics {
332                        metrics.plotting_capacity_used.inc();
333                    }
334
335                    // Give a chance to interrupt plotting if necessary
336                    yield_now().await;
337
338                    if !progress_updater
339                        .update_progress_and_events(
340                            &mut progress_sender,
341                            SectorPlottingProgress::Encoding,
342                        )
343                        .await
344                    {
345                        if let Some(metrics) = &metrics {
346                            metrics.plotting_capacity_used.dec();
347                        }
348                        return;
349                    }
350
351                    let encoding_start = Instant::now();
352
353                    let plotting_result = tokio::task::block_in_place(move || {
354                        let encoded_sector = encode_sector(
355                            downloaded_sector,
356                            EncodeSectorOptions {
357                                sector_index,
358                                records_encoder: &mut *records_encoder,
359                                abort_early: &abort_early,
360                            },
361                        )?;
362
363                        if abort_early.load(Ordering::Acquire) {
364                            return Err(PlottingError::AbortEarly);
365                        }
366
367                        drop(records_encoder);
368
369                        let mut sector = Vec::new();
370
371                        write_sector(&encoded_sector, &mut sector)?;
372
373                        Ok((sector, encoded_sector.plotted_sector))
374                    });
375
376                    if let Some(metrics) = &metrics {
377                        metrics.plotting_capacity_used.dec();
378                    }
379
380                    match plotting_result {
381                        Ok(plotting_result) => {
382                            if !progress_updater
383                                .update_progress_and_events(
384                                    &mut progress_sender,
385                                    SectorPlottingProgress::Encoded(encoding_start.elapsed()),
386                                )
387                                .await
388                            {
389                                return;
390                            }
391
392                            plotting_result
393                        }
394                        Err(PlottingError::AbortEarly) => {
395                            return;
396                        }
397                        Err(error) => {
398                            progress_updater
399                                .update_progress_and_events(
400                                    &mut progress_sender,
401                                    SectorPlottingProgress::Error {
402                                        error: format!("Failed to encode sector: {error}"),
403                                    },
404                                )
405                                .await;
406
407                            return;
408                        }
409                    }
410                };
411
412                progress_updater
413                    .update_progress_and_events(
414                        &mut progress_sender,
415                        SectorPlottingProgress::Finished {
416                            plotted_sector,
417                            time: start.elapsed(),
418                            sector: Box::pin({
419                                let mut sector = Some(Ok(Bytes::from(sector)));
420
421                                stream::poll_fn(move |_cx| {
422                                    // Just so that permit is dropped with stream itself
423                                    let _downloading_permit = &downloading_permit;
424
425                                    Poll::Ready(sector.take())
426                                })
427                            }),
428                        },
429                    )
430                    .await;
431            }
432        };
433
434        // Spawn a separate task such that `block_in_place` inside will not affect anything else
435        let plotting_task =
436            AsyncJoinOnDrop::new(tokio::spawn(plotting_fut.in_current_span()), true);
437        if let Err(error) = self.tasks_sender.clone().send(plotting_task).await {
438            warn!(%error, "Failed to send plotting task");
439
440            let progress = SectorPlottingProgress::Error {
441                error: format!("Failed to send plotting task: {error}"),
442            };
443
444            self.handlers
445                .plotting_progress
446                .call_simple(&public_key, &sector_index, &progress);
447        }
448    }
449}
450
451struct ProgressUpdater {
452    public_key: PublicKey,
453    sector_index: SectorIndex,
454    handlers: Arc<Handlers>,
455    metrics: Option<Arc<GpuPlotterMetrics>>,
456}
457
458impl ProgressUpdater {
459    /// Returns `true` on success and `false` if progress receiver channel is gone
460    async fn update_progress_and_events<PS>(
461        &self,
462        progress_sender: &mut PS,
463        progress: SectorPlottingProgress,
464    ) -> bool
465    where
466        PS: Sink<SectorPlottingProgress> + Unpin,
467        PS::Error: Error,
468    {
469        if let Some(metrics) = &self.metrics {
470            match &progress {
471                SectorPlottingProgress::Downloading => {
472                    metrics.sector_downloading.inc();
473                }
474                SectorPlottingProgress::Downloaded(time) => {
475                    metrics.sector_downloading_time.observe(time.as_secs_f64());
476                    metrics.sector_downloaded.inc();
477                }
478                SectorPlottingProgress::Encoding => {
479                    metrics.sector_encoding.inc();
480                }
481                SectorPlottingProgress::Encoded(time) => {
482                    metrics.sector_encoding_time.observe(time.as_secs_f64());
483                    metrics.sector_encoded.inc();
484                }
485                SectorPlottingProgress::Finished { time, .. } => {
486                    metrics.sector_plotting_time.observe(time.as_secs_f64());
487                    metrics.sector_plotted.inc();
488                }
489                SectorPlottingProgress::Error { .. } => {
490                    metrics.sector_plotting_error.inc();
491                }
492            }
493        }
494        self.handlers.plotting_progress.call_simple(
495            &self.public_key,
496            &self.sector_index,
497            &progress,
498        );
499
500        if let Err(error) = progress_sender.send(progress).await {
501            warn!(%error, "Failed to send progress update");
502
503            false
504        } else {
505            true
506        }
507    }
508}