1use crate::{
4 backend::{Completion, Submission},
5 observation::{ObservationError, TensorObservation, TensorObservationData},
6 scheduler::{
7 RequestId, RequestStatus, Scheduler, SchedulerCapabilities, SchedulerError,
8 SchedulerLimits, SchedulerReport, SemanticStateTransaction, TransitionOutput,
9 WorkDescriptor, WorkId,
10 },
11};
12use serde::{Deserialize, Serialize};
13use std::{collections::BTreeMap, fmt::Debug, time::Instant};
14
15pub const MAX_REALTIME_FRAME_DELAY: usize = i32::MAX as usize;
21
22#[derive(Debug, Clone, Copy, Eq, PartialEq, Serialize, Deserialize)]
24#[serde(rename_all = "snake_case")]
25#[non_exhaustive]
26pub enum RealtimeFrameConvention {
27 FeedbackAlignedHistory,
30 AbsoluteDelayedSlots,
33}
34
35#[derive(Debug, Clone, Eq, PartialEq, Serialize)]
37pub struct RealtimeSpeechConfig {
38 total_audio_codebooks: usize,
39 input_audio_codebooks: usize,
40 generated_audio_codebooks: usize,
41 depth_audio_codebooks: usize,
42 text_padding_token: i32,
43 audio_padding_token: i32,
44 frame_convention: RealtimeFrameConvention,
45 delays: Vec<usize>,
46}
47
48impl<'de> Deserialize<'de> for RealtimeSpeechConfig {
49 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
50 where
51 D: serde::Deserializer<'de>,
52 {
53 #[derive(Deserialize)]
54 struct Raw {
55 total_audio_codebooks: usize,
56 input_audio_codebooks: usize,
57 generated_audio_codebooks: usize,
58 depth_audio_codebooks: usize,
59 text_padding_token: i32,
60 audio_padding_token: i32,
61 frame_convention: RealtimeFrameConvention,
62 delays: Vec<usize>,
63 }
64 let raw = Raw::deserialize(deserializer)?;
65 Self::new(
66 raw.total_audio_codebooks,
67 raw.input_audio_codebooks,
68 raw.generated_audio_codebooks,
69 raw.depth_audio_codebooks,
70 raw.text_padding_token,
71 raw.audio_padding_token,
72 raw.frame_convention,
73 raw.delays,
74 )
75 .map_err(serde::de::Error::custom)
76 }
77}
78
79impl RealtimeSpeechConfig {
80 #[allow(clippy::too_many_arguments)] pub fn new(
83 total_audio_codebooks: usize,
84 input_audio_codebooks: usize,
85 generated_audio_codebooks: usize,
86 depth_audio_codebooks: usize,
87 text_padding_token: i32,
88 audio_padding_token: i32,
89 frame_convention: RealtimeFrameConvention,
90 delays: Vec<usize>,
91 ) -> Result<Self, RealtimeConfigError> {
92 if total_audio_codebooks == 0
93 || generated_audio_codebooks == 0
94 || depth_audio_codebooks == 0
95 {
96 return Err(RealtimeConfigError::EmptyCodebookGeometry);
97 }
98 if input_audio_codebooks.checked_add(generated_audio_codebooks)
99 != Some(total_audio_codebooks)
100 {
101 return Err(RealtimeConfigError::CodebookPartition {
102 total: total_audio_codebooks,
103 input: input_audio_codebooks,
104 generated: generated_audio_codebooks,
105 });
106 }
107 if generated_audio_codebooks > depth_audio_codebooks
108 || depth_audio_codebooks > total_audio_codebooks
109 {
110 return Err(RealtimeConfigError::DepthCodebookGeometry {
111 generated: generated_audio_codebooks,
112 depth: depth_audio_codebooks,
113 total: total_audio_codebooks,
114 });
115 }
116 if text_padding_token < 0 || audio_padding_token < 0 {
117 return Err(RealtimeConfigError::NegativePaddingToken {
118 text: text_padding_token,
119 audio: audio_padding_token,
120 });
121 }
122 let expected_delays = total_audio_codebooks
123 .checked_add(1)
124 .ok_or(RealtimeConfigError::CodebookCountOverflow)?;
125 if delays.len() != expected_delays {
126 return Err(RealtimeConfigError::DelayCount {
127 expected: expected_delays,
128 actual: delays.len(),
129 });
130 }
131 if let Some((slot, delay)) = delays
132 .iter()
133 .copied()
134 .enumerate()
135 .find(|(_, delay)| *delay > MAX_REALTIME_FRAME_DELAY)
136 {
137 return Err(RealtimeConfigError::DelayOutOfRange {
138 slot,
139 delay,
140 maximum: MAX_REALTIME_FRAME_DELAY,
141 });
142 }
143 Ok(Self {
144 total_audio_codebooks,
145 input_audio_codebooks,
146 generated_audio_codebooks,
147 depth_audio_codebooks,
148 text_padding_token,
149 audio_padding_token,
150 frame_convention,
151 delays,
152 })
153 }
154
155 pub const fn total_audio_codebooks(&self) -> usize {
157 self.total_audio_codebooks
158 }
159 pub const fn input_audio_codebooks(&self) -> usize {
161 self.input_audio_codebooks
162 }
163 pub const fn generated_audio_codebooks(&self) -> usize {
165 self.generated_audio_codebooks
166 }
167 pub const fn depth_audio_codebooks(&self) -> usize {
169 self.depth_audio_codebooks
170 }
171 pub const fn text_padding_token(&self) -> i32 {
173 self.text_padding_token
174 }
175 pub const fn audio_padding_token(&self) -> i32 {
177 self.audio_padding_token
178 }
179 pub const fn frame_convention(&self) -> RealtimeFrameConvention {
181 self.frame_convention
182 }
183 pub fn delays(&self) -> &[usize] {
185 &self.delays
186 }
187 pub fn text_delay(&self) -> usize {
189 self.delays[0]
190 }
191 pub fn audio_delays(&self) -> &[usize] {
193 &self.delays[1..]
194 }
195 pub fn max_delay(&self) -> usize {
197 self.delays.iter().copied().max().unwrap_or(0)
198 }
199 pub fn max_audio_delay(&self) -> usize {
201 self.audio_delays().iter().copied().max().unwrap_or(0)
202 }
203}
204
205#[derive(Debug, Clone, PartialEq, thiserror::Error)]
207#[non_exhaustive]
208pub enum RealtimeConfigError {
209 #[error("realtime codebook geometry must be nonzero")]
211 EmptyCodebookGeometry,
212 #[error(
214 "realtime input ({input}) and generated ({generated}) codebooks do not partition total {total}"
215 )]
216 CodebookPartition {
217 total: usize,
219 input: usize,
221 generated: usize,
223 },
224 #[error(
227 "realtime generated ({generated}), depth ({depth}), and total ({total}) codebooks must satisfy 0 < generated <= depth <= total"
228 )]
229 DepthCodebookGeometry {
230 generated: usize,
232 depth: usize,
234 total: usize,
236 },
237 #[error("realtime text-plus-audio slot count overflowed")]
239 CodebookCountOverflow,
240 #[error("realtime padding tokens must be non-negative, got text={text} audio={audio}")]
242 NegativePaddingToken {
243 text: i32,
245 audio: i32,
247 },
248 #[error("realtime delay schedule has {actual} entries, expected {expected}")]
250 DelayCount {
251 expected: usize,
253 actual: usize,
255 },
256 #[error("realtime delay {delay} at slot {slot} exceeds maximum {maximum}")]
258 DelayOutOfRange {
259 slot: usize,
261 delay: usize,
263 maximum: usize,
265 },
266 #[error(
268 "realtime sampling temperatures must be finite and non-negative, got text={text} audio={audio}"
269 )]
270 SamplingTemperature {
271 text: f32,
273 audio: f32,
275 },
276 #[error(
278 "realtime sampling top-k must be positive when set, got text={text:?} audio={audio:?}"
279 )]
280 SamplingTopK {
281 text: Option<usize>,
283 audio: Option<usize>,
285 },
286}
287
288#[derive(Debug, Clone, Copy, Eq, Hash, Ord, PartialEq, PartialOrd)]
290#[non_exhaustive]
291pub enum RealtimeFrameSlot {
292 Text,
294 Audio(usize),
296}
297
298impl RealtimeFrameSlot {
299 fn index(self) -> usize {
300 match self {
301 Self::Text => 0,
302 Self::Audio(codebook) => codebook + 1,
303 }
304 }
305}
306
307#[derive(Debug, Clone, Copy, Eq, Hash, Ord, PartialEq, PartialOrd)]
309pub struct RealtimeSlotCoordinate {
310 position: usize,
311 slot: RealtimeFrameSlot,
312}
313
314impl RealtimeSlotCoordinate {
315 pub const fn new(position: usize, slot: RealtimeFrameSlot) -> Self {
317 Self { position, slot }
318 }
319
320 pub const fn position(self) -> usize {
322 self.position
323 }
324
325 pub const fn slot(self) -> RealtimeFrameSlot {
327 self.slot
328 }
329}
330
331#[derive(Debug, Clone, Copy, Eq, PartialEq)]
335#[non_exhaustive]
336pub enum RealtimeSlotOccupancy {
337 Input,
339 Forced,
341 Padding,
343 Sampled,
345}
346
347#[derive(Debug, Clone, Copy, Eq, PartialEq)]
349#[non_exhaustive]
350pub enum RealtimeTemporalSource {
351 Padding(RealtimeFrameSlot),
353 Occupied {
355 coordinate: RealtimeSlotCoordinate,
357 occupancy: RealtimeSlotOccupancy,
359 },
360}
361
362#[derive(Debug, Clone, Copy, Eq, PartialEq)]
364#[non_exhaustive]
365pub enum RealtimeTargetSource {
366 Forced,
368 Sampled,
370 Existing(RealtimeSlotOccupancy),
372}
373
374#[derive(Debug, Clone, Copy, Eq, PartialEq)]
376pub struct RealtimeTargetDecision {
377 slot: RealtimeFrameSlot,
378 coordinate: Option<RealtimeSlotCoordinate>,
379 source: RealtimeTargetSource,
380}
381
382impl RealtimeTargetDecision {
383 pub const fn slot(self) -> RealtimeFrameSlot {
385 self.slot
386 }
387
388 pub const fn coordinate(self) -> Option<RealtimeSlotCoordinate> {
390 self.coordinate
391 }
392
393 pub const fn source(self) -> RealtimeTargetSource {
395 self.source
396 }
397}
398
399#[derive(Debug, Clone, Eq, PartialEq)]
401pub struct RealtimeFrameForcing {
402 text: bool,
403 generated_audio: Vec<bool>,
404}
405
406impl RealtimeFrameForcing {
407 pub fn new(text: bool, generated_audio: Vec<bool>) -> Self {
409 Self {
410 text,
411 generated_audio,
412 }
413 }
414
415 pub fn none(config: &RealtimeSpeechConfig) -> Self {
417 Self::new(false, vec![false; config.generated_audio_codebooks])
418 }
419
420 pub const fn text(&self) -> bool {
422 self.text
423 }
424
425 pub fn generated_audio(&self) -> &[bool] {
427 &self.generated_audio
428 }
429}
430
431#[derive(Debug, Clone, Eq, PartialEq)]
433pub struct RealtimeFrameTransition {
434 frontier: usize,
435 input_placements: Vec<RealtimeSlotCoordinate>,
436 forced_placements: Vec<RealtimeSlotCoordinate>,
437 warmup_padding: Vec<RealtimeSlotCoordinate>,
438 temporal_inputs: Vec<RealtimeTemporalSource>,
439 targets: Vec<RealtimeTargetDecision>,
440 output: Option<Vec<RealtimeSlotCoordinate>>,
441 model_call_required: bool,
442 next_frontier: usize,
443}
444
445impl RealtimeFrameTransition {
446 pub const fn frontier(&self) -> usize {
448 self.frontier
449 }
450
451 pub fn input_placements(&self) -> &[RealtimeSlotCoordinate] {
453 &self.input_placements
454 }
455
456 pub fn forced_placements(&self) -> &[RealtimeSlotCoordinate] {
458 &self.forced_placements
459 }
460
461 pub fn warmup_padding(&self) -> &[RealtimeSlotCoordinate] {
463 &self.warmup_padding
464 }
465
466 pub fn temporal_inputs(&self) -> &[RealtimeTemporalSource] {
468 &self.temporal_inputs
469 }
470
471 pub fn targets(&self) -> &[RealtimeTargetDecision] {
473 &self.targets
474 }
475
476 pub fn output(&self) -> Option<&[RealtimeSlotCoordinate]> {
478 self.output.as_deref()
479 }
480
481 pub const fn model_call_required(&self) -> bool {
483 self.model_call_required
484 }
485
486 pub const fn next_frontier(&self) -> usize {
488 self.next_frontier
489 }
490}
491
492#[derive(Debug, Clone, Eq, PartialEq)]
494pub struct RealtimeFrameScheduleState {
495 schedule: RealtimeSpeechConfig,
496 frontier: usize,
497 occupied: BTreeMap<RealtimeSlotCoordinate, RealtimeSlotOccupancy>,
498}
499
500impl RealtimeFrameScheduleState {
501 pub fn new(schedule: RealtimeSpeechConfig) -> Self {
503 Self {
504 schedule,
505 frontier: 0,
506 occupied: BTreeMap::new(),
507 }
508 }
509
510 pub const fn schedule(&self) -> &RealtimeSpeechConfig {
512 &self.schedule
513 }
514
515 pub const fn frontier(&self) -> usize {
517 self.frontier
518 }
519
520 pub fn occupancy(&self, coordinate: RealtimeSlotCoordinate) -> Option<RealtimeSlotOccupancy> {
522 self.occupied.get(&coordinate).copied()
523 }
524
525 pub fn validate_schedule(
527 &self,
528 schedule: &RealtimeSpeechConfig,
529 ) -> Result<(), RealtimeScheduleError> {
530 if &self.schedule == schedule {
531 Ok(())
532 } else {
533 Err(RealtimeScheduleError::ScheduleMismatch)
534 }
535 }
536
537 pub fn advance(
543 &mut self,
544 schedule: &RealtimeSpeechConfig,
545 forcing: &RealtimeFrameForcing,
546 ) -> Result<RealtimeFrameTransition, RealtimeScheduleError> {
547 self.validate_schedule(schedule)?;
548 if forcing.generated_audio.len() != schedule.generated_audio_codebooks {
549 return Err(RealtimeScheduleError::ForcingCount {
550 expected: schedule.generated_audio_codebooks,
551 actual: forcing.generated_audio.len(),
552 });
553 }
554 let mut branch = self.clone();
555 let transition = match schedule.frame_convention {
556 RealtimeFrameConvention::FeedbackAlignedHistory => branch.advance_feedback(forcing)?,
557 RealtimeFrameConvention::AbsoluteDelayedSlots => branch.advance_absolute(forcing)?,
558 };
559 *self = branch;
560 Ok(transition)
561 }
562
563 fn advance_feedback(
564 &mut self,
565 forcing: &RealtimeFrameForcing,
566 ) -> Result<RealtimeFrameTransition, RealtimeScheduleError> {
567 let schedule = &self.schedule;
568 let frontier = self.frontier;
569 let next_frontier = checked_add(frontier, 1)?;
570 let generated = schedule.generated_audio_codebooks;
571 let mut input_placements = Vec::with_capacity(schedule.input_audio_codebooks);
572 for codebook in generated..schedule.total_audio_codebooks {
573 let coordinate = coordinate(frontier, RealtimeFrameSlot::Audio(codebook));
574 self.occupied
575 .insert(coordinate, RealtimeSlotOccupancy::Input);
576 input_placements.push(coordinate);
577 }
578
579 let mut forced_placements = Vec::new();
580 for (codebook, forced) in forcing.generated_audio.iter().copied().enumerate() {
581 if forced {
582 let coordinate = coordinate(frontier, RealtimeFrameSlot::Audio(codebook));
583 self.occupied
584 .insert(coordinate, RealtimeSlotOccupancy::Forced);
585 forced_placements.push(coordinate);
586 }
587 }
588
589 let mut temporal_inputs = Vec::with_capacity(schedule.delays.len());
590 for slot in slots(schedule.total_audio_codebooks) {
591 let delay = schedule.delays[slot.index()];
592 let source = frontier
593 .checked_sub(1)
594 .and_then(|position| position.checked_sub(delay));
595 match source {
596 None => temporal_inputs.push(RealtimeTemporalSource::Padding(slot)),
597 Some(position) => {
598 let coordinate = coordinate(position, slot);
599 let occupancy = self.required_occupancy(coordinate)?;
600 temporal_inputs.push(RealtimeTemporalSource::Occupied {
601 coordinate,
602 occupancy,
603 });
604 }
605 }
606 }
607
608 let mut targets = Vec::with_capacity(1 + schedule.depth_audio_codebooks);
609 let text_coordinate = frontier
610 .checked_sub(schedule.text_delay())
611 .map(|position| coordinate(position, RealtimeFrameSlot::Text));
612 let text_source = if forcing.text {
613 RealtimeTargetSource::Forced
614 } else {
615 RealtimeTargetSource::Sampled
616 };
617 if let Some(coordinate) = text_coordinate {
618 self.occupied.insert(
619 coordinate,
620 if forcing.text {
621 RealtimeSlotOccupancy::Forced
622 } else {
623 RealtimeSlotOccupancy::Sampled
624 },
625 );
626 if forcing.text {
627 forced_placements.push(coordinate);
628 }
629 }
630 targets.push(RealtimeTargetDecision {
631 slot: RealtimeFrameSlot::Text,
632 coordinate: text_coordinate,
633 source: text_source,
634 });
635 for codebook in 0..schedule.depth_audio_codebooks {
636 let slot = RealtimeFrameSlot::Audio(codebook);
637 if codebook < generated {
638 let target_coordinate = frontier
639 .checked_sub(schedule.audio_delays()[codebook])
640 .map(|position| coordinate(position, slot));
641 let forced = forcing.generated_audio[codebook];
642 if let Some(coordinate) = target_coordinate {
643 self.occupied.insert(
644 coordinate,
645 if forced {
646 RealtimeSlotOccupancy::Forced
647 } else {
648 RealtimeSlotOccupancy::Sampled
649 },
650 );
651 }
652 targets.push(RealtimeTargetDecision {
653 slot,
654 coordinate: target_coordinate,
655 source: if forced {
656 RealtimeTargetSource::Forced
657 } else {
658 RealtimeTargetSource::Sampled
659 },
660 });
661 } else {
662 let input_coordinate = coordinate(frontier, slot);
663 let occupancy = self.required_occupancy(input_coordinate)?;
664 targets.push(RealtimeTargetDecision {
665 slot,
666 coordinate: Some(input_coordinate),
667 source: RealtimeTargetSource::Existing(occupancy),
668 });
669 }
670 }
671
672 let output = frontier
673 .checked_sub(schedule.max_delay())
674 .map(|position| self.output_at_same_position(position))
675 .transpose()?;
676 self.frontier = next_frontier;
677 self.prune_before(frontier.saturating_sub(schedule.max_delay()));
678 Ok(RealtimeFrameTransition {
679 frontier,
680 input_placements,
681 forced_placements,
682 warmup_padding: Vec::new(),
683 temporal_inputs,
684 targets,
685 output,
686 model_call_required: true,
687 next_frontier,
688 })
689 }
690
691 fn advance_absolute(
692 &mut self,
693 forcing: &RealtimeFrameForcing,
694 ) -> Result<RealtimeFrameTransition, RealtimeScheduleError> {
695 let schedule = &self.schedule;
696 let frontier = self.frontier;
697 let next_frontier = checked_add(frontier, 1)?;
698 let generated = schedule.generated_audio_codebooks;
699 let mut input_placements = Vec::with_capacity(schedule.input_audio_codebooks);
700 for codebook in generated..schedule.total_audio_codebooks {
701 let position = checked_add(frontier, schedule.audio_delays()[codebook])?;
702 let coordinate = coordinate(position, RealtimeFrameSlot::Audio(codebook));
703 self.occupied
704 .insert(coordinate, RealtimeSlotOccupancy::Input);
705 input_placements.push(coordinate);
706 }
707
708 let mut forced_placements = Vec::new();
709 if forcing.text {
710 let position = checked_add(frontier, schedule.text_delay())?;
711 let coordinate = coordinate(position, RealtimeFrameSlot::Text);
712 self.occupied
713 .insert(coordinate, RealtimeSlotOccupancy::Forced);
714 forced_placements.push(coordinate);
715 }
716 for (codebook, forced) in forcing.generated_audio.iter().copied().enumerate() {
717 if forced {
718 let position = checked_add(frontier, schedule.audio_delays()[codebook])?;
719 let coordinate = coordinate(position, RealtimeFrameSlot::Audio(codebook));
720 self.occupied
721 .insert(coordinate, RealtimeSlotOccupancy::Forced);
722 forced_placements.push(coordinate);
723 }
724 }
725
726 let mut warmup_padding = Vec::new();
727 for slot in slots(schedule.total_audio_codebooks) {
728 if frontier <= schedule.delays[slot.index()] {
729 let coordinate = coordinate(frontier, slot);
730 self.occupied
731 .insert(coordinate, RealtimeSlotOccupancy::Padding);
732 warmup_padding.push(coordinate);
733 }
734 }
735
736 let mut temporal_inputs = Vec::new();
737 let mut targets = Vec::new();
738 let output = if frontier == 0 {
739 None
740 } else {
741 let input_position = frontier - 1;
742 temporal_inputs.reserve(schedule.delays.len());
743 for slot in slots(schedule.total_audio_codebooks) {
744 let coordinate = coordinate(input_position, slot);
745 let occupancy = self.required_occupancy(coordinate)?;
746 temporal_inputs.push(RealtimeTemporalSource::Occupied {
747 coordinate,
748 occupancy,
749 });
750 }
751 targets.reserve(1 + schedule.depth_audio_codebooks);
752 for slot in std::iter::once(RealtimeFrameSlot::Text)
753 .chain((0..schedule.depth_audio_codebooks).map(RealtimeFrameSlot::Audio))
754 {
755 let coordinate = coordinate(frontier, slot);
756 let (source, occupancy) = match self.occupied.get(&coordinate).copied() {
757 Some(RealtimeSlotOccupancy::Forced) => {
758 (RealtimeTargetSource::Forced, RealtimeSlotOccupancy::Forced)
759 }
760 Some(occupancy) => (RealtimeTargetSource::Existing(occupancy), occupancy),
761 None => (
762 RealtimeTargetSource::Sampled,
763 RealtimeSlotOccupancy::Sampled,
764 ),
765 };
766 self.occupied.insert(coordinate, occupancy);
767 targets.push(RealtimeTargetDecision {
768 slot,
769 coordinate: Some(coordinate),
770 source,
771 });
772 }
773 if frontier <= schedule.max_delay() {
774 None
775 } else {
776 let base = frontier - schedule.max_delay();
777 let coordinates = (0..generated)
778 .map(|codebook| {
779 let position = checked_add(base, schedule.audio_delays()[codebook])?;
780 let coordinate = coordinate(position, RealtimeFrameSlot::Audio(codebook));
781 self.required_occupancy(coordinate)?;
782 Ok(coordinate)
783 })
784 .collect::<Result<Vec<_>, RealtimeScheduleError>>()?;
785 Some(coordinates)
786 }
787 };
788
789 self.frontier = next_frontier;
790 self.prune_before(next_frontier.saturating_sub(schedule.max_delay().saturating_add(1)));
791 Ok(RealtimeFrameTransition {
792 frontier,
793 input_placements,
794 forced_placements,
795 warmup_padding,
796 temporal_inputs,
797 targets,
798 output,
799 model_call_required: frontier != 0,
800 next_frontier,
801 })
802 }
803
804 fn required_occupancy(
805 &self,
806 coordinate: RealtimeSlotCoordinate,
807 ) -> Result<RealtimeSlotOccupancy, RealtimeScheduleError> {
808 self.occupied
809 .get(&coordinate)
810 .copied()
811 .ok_or(RealtimeScheduleError::MissingSlot { coordinate })
812 }
813
814 fn output_at_same_position(
815 &self,
816 position: usize,
817 ) -> Result<Vec<RealtimeSlotCoordinate>, RealtimeScheduleError> {
818 (0..self.schedule.generated_audio_codebooks)
819 .map(|codebook| {
820 let coordinate = coordinate(position, RealtimeFrameSlot::Audio(codebook));
821 self.required_occupancy(coordinate)?;
822 Ok(coordinate)
823 })
824 .collect()
825 }
826
827 fn prune_before(&mut self, minimum: usize) {
828 self.occupied
829 .retain(|coordinate, _| coordinate.position >= minimum);
830 }
831}
832
833impl SemanticStateTransaction for RealtimeFrameScheduleState {
834 type Branch = Self;
835 type Error = RealtimeScheduleError;
836
837 fn branch(&self) -> Result<Self::Branch, Self::Error> {
838 Ok(self.clone())
839 }
840
841 fn commit_branch(&mut self, branch: Self::Branch) -> Result<(), Self::Error> {
842 self.validate_schedule(&branch.schedule)?;
843 *self = branch;
844 Ok(())
845 }
846}
847
848#[derive(Debug, Clone, Eq, PartialEq, thiserror::Error)]
850#[non_exhaustive]
851pub enum RealtimeScheduleError {
852 #[error("realtime frame schedule state does not match the normalized schedule")]
854 ScheduleMismatch,
855 #[error("realtime forcing mask has {actual} audio entries, expected {expected}")]
857 ForcingCount {
858 expected: usize,
860 actual: usize,
862 },
863 #[error("realtime delayed slot {coordinate:?} is not occupied")]
865 MissingSlot {
866 coordinate: RealtimeSlotCoordinate,
868 },
869 #[error("realtime frame coordinate overflowed")]
871 CoordinateOverflow,
872}
873
874fn checked_add(left: usize, right: usize) -> Result<usize, RealtimeScheduleError> {
875 left.checked_add(right)
876 .ok_or(RealtimeScheduleError::CoordinateOverflow)
877}
878
879fn coordinate(position: usize, slot: RealtimeFrameSlot) -> RealtimeSlotCoordinate {
880 RealtimeSlotCoordinate::new(position, slot)
881}
882
883fn slots(total_audio_codebooks: usize) -> impl Iterator<Item = RealtimeFrameSlot> {
884 std::iter::once(RealtimeFrameSlot::Text)
885 .chain((0..total_audio_codebooks).map(RealtimeFrameSlot::Audio))
886}
887
888#[derive(Debug, Clone, Copy, PartialEq, Serialize)]
890pub struct RealtimeSampling {
891 text_temperature: f32,
892 audio_temperature: f32,
893 text_top_k: Option<usize>,
894 audio_top_k: Option<usize>,
895 seed: u64,
896}
897
898impl<'de> Deserialize<'de> for RealtimeSampling {
899 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
900 where
901 D: serde::Deserializer<'de>,
902 {
903 #[derive(Deserialize)]
904 struct Raw {
905 text_temperature: f32,
906 audio_temperature: f32,
907 text_top_k: Option<usize>,
908 audio_top_k: Option<usize>,
909 seed: u64,
910 }
911 let raw = Raw::deserialize(deserializer)?;
912 Self::new(raw.text_temperature, raw.audio_temperature, raw.seed)
913 .and_then(|sampling| sampling.with_top_k(raw.text_top_k, raw.audio_top_k))
914 .map_err(serde::de::Error::custom)
915 }
916}
917
918impl RealtimeSampling {
919 pub fn new(
921 text_temperature: f32,
922 audio_temperature: f32,
923 seed: u64,
924 ) -> Result<Self, RealtimeConfigError> {
925 if !text_temperature.is_finite()
926 || text_temperature < 0.0
927 || !audio_temperature.is_finite()
928 || audio_temperature < 0.0
929 {
930 return Err(RealtimeConfigError::SamplingTemperature {
931 text: text_temperature,
932 audio: audio_temperature,
933 });
934 }
935 Ok(Self {
936 text_temperature,
937 audio_temperature,
938 text_top_k: None,
939 audio_top_k: None,
940 seed,
941 })
942 }
943
944 pub fn with_top_k(
946 mut self,
947 text_top_k: Option<usize>,
948 audio_top_k: Option<usize>,
949 ) -> Result<Self, RealtimeConfigError> {
950 if text_top_k == Some(0) || audio_top_k == Some(0) {
951 return Err(RealtimeConfigError::SamplingTopK {
952 text: text_top_k,
953 audio: audio_top_k,
954 });
955 }
956 self.text_top_k = text_top_k;
957 self.audio_top_k = audio_top_k;
958 Ok(self)
959 }
960
961 pub const fn greedy() -> Self {
963 Self {
964 text_temperature: 0.0,
965 audio_temperature: 0.0,
966 text_top_k: None,
967 audio_top_k: None,
968 seed: 0,
969 }
970 }
971 pub const fn text_temperature(self) -> f32 {
973 self.text_temperature
974 }
975 pub const fn audio_temperature(self) -> f32 {
977 self.audio_temperature
978 }
979 pub const fn text_top_k(self) -> Option<usize> {
981 self.text_top_k
982 }
983 pub const fn audio_top_k(self) -> Option<usize> {
985 self.audio_top_k
986 }
987 pub const fn seed(self) -> u64 {
989 self.seed
990 }
991 pub const fn is_stochastic(self) -> bool {
993 self.text_temperature != 0.0 || self.audio_temperature != 0.0
994 }
995}
996
997impl Default for RealtimeSampling {
998 fn default() -> Self {
999 Self::greedy()
1000 }
1001}
1002
1003#[derive(Debug, Clone, Eq, PartialEq)]
1005pub struct RealtimeInputFrame {
1006 batch: usize,
1007 input_audio_tokens: Vec<i32>,
1008 forced_generated_audio_tokens: Option<Vec<i32>>,
1009 forced_generated_audio_codebooks: Option<Vec<bool>>,
1010 forced_text_tokens: Option<Vec<i32>>,
1011 retain_diagnostics: bool,
1012}
1013
1014impl RealtimeInputFrame {
1015 pub fn new(batch: usize, input_audio_tokens: Vec<i32>) -> Self {
1017 Self {
1018 batch,
1019 input_audio_tokens,
1020 forced_generated_audio_tokens: None,
1021 forced_generated_audio_codebooks: None,
1022 forced_text_tokens: None,
1023 retain_diagnostics: false,
1024 }
1025 }
1026
1027 pub fn with_forced_generated_audio(mut self, tokens: Vec<i32>) -> Self {
1029 self.forced_generated_audio_tokens = Some(tokens);
1030 self.forced_generated_audio_codebooks = None;
1031 self
1032 }
1033
1034 pub fn with_partially_forced_generated_audio(
1036 mut self,
1037 tokens: Vec<i32>,
1038 codebooks: Vec<bool>,
1039 ) -> Self {
1040 self.forced_generated_audio_tokens = Some(tokens);
1041 self.forced_generated_audio_codebooks = Some(codebooks);
1042 self
1043 }
1044
1045 pub fn with_forced_text(mut self, tokens: Vec<i32>) -> Self {
1047 self.forced_text_tokens = Some(tokens);
1048 self
1049 }
1050
1051 pub fn with_diagnostics(mut self) -> Self {
1053 self.retain_diagnostics = true;
1054 self
1055 }
1056
1057 pub const fn batch(&self) -> usize {
1059 self.batch
1060 }
1061 pub fn input_audio_tokens(&self) -> &[i32] {
1063 &self.input_audio_tokens
1064 }
1065 pub fn forced_generated_audio_tokens(&self) -> Option<&[i32]> {
1067 self.forced_generated_audio_tokens.as_deref()
1068 }
1069 pub fn forced_generated_audio_codebooks(&self) -> Option<&[bool]> {
1071 self.forced_generated_audio_codebooks.as_deref()
1072 }
1073 pub fn forced_text_tokens(&self) -> Option<&[i32]> {
1075 self.forced_text_tokens.as_deref()
1076 }
1077 pub const fn retains_diagnostics(&self) -> bool {
1079 self.retain_diagnostics
1080 }
1081}
1082
1083#[derive(Debug, Clone, PartialEq)]
1085pub struct RealtimeDecisionDiagnostics {
1086 prediction: usize,
1087 tensor: TensorObservation,
1088}
1089
1090impl RealtimeDecisionDiagnostics {
1091 pub fn new(
1093 prediction: usize,
1094 shape: Vec<usize>,
1095 logits: Vec<f32>,
1096 ) -> Result<Self, ObservationError> {
1097 Ok(Self {
1098 prediction,
1099 tensor: TensorObservation::new(shape, TensorObservationData::F32(logits))?,
1100 })
1101 }
1102 pub const fn prediction(&self) -> usize {
1104 self.prediction
1105 }
1106 pub fn shape(&self) -> &[usize] {
1108 self.tensor.shape()
1109 }
1110 pub fn logits(&self) -> &[f32] {
1112 let TensorObservationData::F32(values) = self.tensor.data() else {
1113 unreachable!("realtime diagnostics are constructed from F32 values")
1114 };
1115 values
1116 }
1117
1118 pub const fn tensor(&self) -> &TensorObservation {
1120 &self.tensor
1121 }
1122}
1123
1124#[derive(Debug, Clone, PartialEq)]
1126pub struct RealtimeOutputFrame {
1127 batch: usize,
1128 text_tokens: Vec<i32>,
1129 decision_audio_tokens: Vec<i32>,
1130 sampled_audio_tokens: Vec<i32>,
1131 output_audio_tokens: Option<Vec<i32>>,
1132 diagnostics: Vec<RealtimeDecisionDiagnostics>,
1133}
1134
1135impl RealtimeOutputFrame {
1136 pub fn new(
1138 batch: usize,
1139 text_tokens: Vec<i32>,
1140 decision_audio_tokens: Vec<i32>,
1141 sampled_audio_tokens: Vec<i32>,
1142 output_audio_tokens: Option<Vec<i32>>,
1143 diagnostics: Vec<RealtimeDecisionDiagnostics>,
1144 ) -> Self {
1145 Self {
1146 batch,
1147 text_tokens,
1148 decision_audio_tokens,
1149 sampled_audio_tokens,
1150 output_audio_tokens,
1151 diagnostics,
1152 }
1153 }
1154 pub const fn batch(&self) -> usize {
1156 self.batch
1157 }
1158 pub fn text_tokens(&self) -> &[i32] {
1160 &self.text_tokens
1161 }
1162 pub fn decision_audio_tokens(&self) -> &[i32] {
1164 &self.decision_audio_tokens
1165 }
1166 pub fn sampled_audio_tokens(&self) -> &[i32] {
1168 &self.sampled_audio_tokens
1169 }
1170 pub fn output_audio_tokens(&self) -> Option<&[i32]> {
1172 self.output_audio_tokens.as_deref()
1173 }
1174 pub fn diagnostics(&self) -> &[RealtimeDecisionDiagnostics] {
1176 &self.diagnostics
1177 }
1178}
1179
1180pub trait RealtimeBackend {
1186 type Model;
1188 type ModelIdentity: Clone + Debug + Eq;
1190 type Input: WorkDescriptor;
1192 type Output;
1194 type Session: SemanticStateTransaction<Error = Self::Error>;
1196 type Completion: Completion<Error = Self::Error>;
1198 type Error: std::error::Error + Send + Sync + 'static;
1200
1201 fn name(&self) -> &str;
1203 fn model_identity(&self, model: &Self::Model) -> Self::ModelIdentity;
1205 fn session_capabilities(&self, model: &Self::Model) -> crate::SessionCapabilities;
1207 fn model_identity_mismatch(
1209 &self,
1210 expected: &Self::ModelIdentity,
1211 actual: &Self::ModelIdentity,
1212 ) -> Option<String> {
1213 (expected != actual).then(|| "model identity".into())
1214 }
1215 fn speech_config(&self, model: &Self::Model) -> RealtimeSpeechConfig;
1217 fn materialize_input(
1219 &self,
1220 model: &Self::Model,
1221 frame: &RealtimeInputFrame,
1222 ) -> Result<Self::Input, Self::Error>;
1223 fn observe_output(&self, output: &Self::Output) -> Result<RealtimeOutputFrame, Self::Error>;
1225 fn create_session(
1227 &self,
1228 model: &Self::Model,
1229 sampling: RealtimeSampling,
1230 ) -> Result<Self::Session, Self::Error>;
1231 fn validate_session(
1233 &self,
1234 model: &Self::Model,
1235 session: &Self::Session,
1236 ) -> Result<(), Self::Error>;
1237 fn validate_input(&self, model: &Self::Model, input: &Self::Input) -> Result<(), Self::Error>;
1239 fn input_batch_size(&self, input: &Self::Input) -> usize;
1241 fn set_sampling(
1243 &self,
1244 session: &mut Self::Session,
1245 sampling: RealtimeSampling,
1246 ) -> Result<(), Self::Error>;
1247 fn submit_step(
1249 &self,
1250 model: &mut Self::Model,
1251 session: &mut <Self::Session as SemanticStateTransaction>::Branch,
1252 input: &Self::Input,
1253 ) -> Result<Submission<Self::Output, Self::Completion>, Self::Error>;
1254 fn retained_resources(&self, _completion: &Self::Completion) -> usize {
1256 0
1257 }
1258}
1259
1260pub trait RealtimeModelLoadingBackend: RealtimeBackend + Sized {
1266 type Preparation;
1268 type LoadOptions;
1270
1271 fn materialize_realtime_model(
1273 &self,
1274 preparation: Self::Preparation,
1275 options: Self::LoadOptions,
1276 ) -> Result<Self::Model, Self::Error>;
1277}
1278
1279pub fn load_realtime_model<B>(
1281 backend: B,
1282 preparation: B::Preparation,
1283) -> Result<RealtimeModel<B>, B::Error>
1284where
1285 B: RealtimeModelLoadingBackend,
1286 B::LoadOptions: Default,
1287{
1288 load_realtime_model_with_options(backend, preparation, B::LoadOptions::default())
1289}
1290
1291pub fn load_realtime_model_with_options<B: RealtimeModelLoadingBackend>(
1293 backend: B,
1294 preparation: B::Preparation,
1295 options: B::LoadOptions,
1296) -> Result<RealtimeModel<B>, B::Error> {
1297 let model = backend.materialize_realtime_model(preparation, options)?;
1298 Ok(RealtimeModel::new(backend, model))
1299}
1300
1301pub struct RealtimeModel<B: RealtimeBackend> {
1303 backend: B,
1304 model: B::Model,
1305}
1306
1307impl<B: RealtimeBackend> RealtimeModel<B> {
1308 pub const fn new(backend: B, model: B::Model) -> Self {
1310 Self { backend, model }
1311 }
1312 pub const fn backend(&self) -> &B {
1314 &self.backend
1315 }
1316 pub const fn model(&self) -> &B::Model {
1318 &self.model
1319 }
1320 pub fn model_mut(&mut self) -> &mut B::Model {
1322 &mut self.model
1323 }
1324 pub fn speech_config(&self) -> RealtimeSpeechConfig {
1326 self.backend.speech_config(&self.model)
1327 }
1328 pub fn session_capabilities(&self) -> crate::SessionCapabilities {
1330 self.backend.session_capabilities(&self.model)
1331 }
1332 pub fn into_parts(self) -> (B, B::Model) {
1334 (self.backend, self.model)
1335 }
1336}
1337
1338pub struct RealtimeSession<B: RealtimeBackend> {
1340 model_identity: B::ModelIdentity,
1341 state: B::Session,
1342 batch_size: Option<usize>,
1343}
1344
1345pub struct RealtimeSessionBranch<B: RealtimeBackend> {
1347 state: <B::Session as SemanticStateTransaction>::Branch,
1348 batch_size: Option<usize>,
1349}
1350
1351impl<B: RealtimeBackend> RealtimeSession<B> {
1352 pub const fn state(&self) -> &B::Session {
1354 &self.state
1355 }
1356 pub fn state_mut(&mut self) -> &mut B::Session {
1358 &mut self.state
1359 }
1360 pub const fn batch_size(&self) -> Option<usize> {
1362 self.batch_size
1363 }
1364}
1365
1366impl<B: RealtimeBackend> SemanticStateTransaction for RealtimeSession<B> {
1367 type Branch = RealtimeSessionBranch<B>;
1368 type Error = B::Error;
1369
1370 fn branch(&self) -> Result<Self::Branch, Self::Error> {
1371 Ok(RealtimeSessionBranch {
1372 state: self.state.branch()?,
1373 batch_size: self.batch_size,
1374 })
1375 }
1376
1377 fn commit_branch(&mut self, branch: Self::Branch) -> Result<(), Self::Error> {
1378 self.state.commit_branch(branch.state)?;
1379 self.batch_size = branch.batch_size;
1380 Ok(())
1381 }
1382
1383 fn discard_branch(branch: Self::Branch) -> Result<(), Self::Error> {
1384 B::Session::discard_branch(branch.state)
1385 }
1386}
1387
1388struct RealtimeTransition<B: RealtimeBackend> {
1389 backend_name: String,
1390 retained_resources: usize,
1391 output: B::Output,
1392 completion: B::Completion,
1393}
1394
1395impl<B: RealtimeBackend> TransitionOutput for RealtimeTransition<B> {
1396 type Error = B::Error;
1397
1398 fn is_complete(&self) -> Result<bool, Self::Error> {
1399 self.completion.is_complete()
1400 }
1401 fn backend_name(&self) -> Option<String> {
1402 Some(self.backend_name.clone())
1403 }
1404 fn retained_resources(&self) -> usize {
1405 self.retained_resources
1406 }
1407}
1408
1409pub struct RealtimeCompletedStep<O> {
1411 work: WorkId,
1412 output: O,
1413}
1414
1415impl<O> RealtimeCompletedStep<O> {
1416 pub const fn work(&self) -> WorkId {
1418 self.work
1419 }
1420 pub const fn output(&self) -> &O {
1422 &self.output
1423 }
1424 pub fn into_parts(self) -> (WorkId, O) {
1426 (self.work, self.output)
1427 }
1428}
1429
1430#[derive(Debug, thiserror::Error)]
1432#[non_exhaustive]
1433pub enum RealtimeError<E: std::error::Error + 'static> {
1434 #[error("realtime backend failed: {0}")]
1436 Backend(#[source] E),
1437 #[error(transparent)]
1439 Scheduler(#[from] SchedulerError),
1440 #[error("realtime model {component} does not match the scheduler model")]
1442 ModelMismatch {
1443 component: String,
1445 },
1446 #[error("realtime request {request} changed batch size from {expected} to {actual}")]
1448 BatchSize {
1449 request: u64,
1451 expected: usize,
1453 actual: usize,
1455 },
1456 #[error("realtime scheduler frame bound must be positive")]
1458 EmptyRunBound,
1459 #[error("realtime request {request} has {queued} queued frames; drain or cancel them before changing sampling")]
1461 SamplingWhileQueued {
1462 request: u64,
1464 queued: usize,
1466 },
1467 #[error("realtime work {work:?} failed asynchronously: {message}")]
1469 Asynchronous {
1470 work: WorkId,
1472 message: String,
1474 },
1475}
1476
1477pub struct RealtimeScheduler<B: RealtimeBackend> {
1479 model_identity: B::ModelIdentity,
1480 scheduler: Scheduler<B::Input, RealtimeSession<B>, RealtimeTransition<B>>,
1481}
1482
1483impl<B: RealtimeBackend> RealtimeScheduler<B> {
1484 pub fn new(
1486 model: &RealtimeModel<B>,
1487 limits: SchedulerLimits,
1488 ) -> Result<Self, RealtimeError<B::Error>> {
1489 Ok(Self {
1490 model_identity: model.backend.model_identity(&model.model),
1491 scheduler: Scheduler::new(limits)?,
1492 })
1493 }
1494
1495 fn validate_model(&self, model: &RealtimeModel<B>) -> Result<(), RealtimeError<B::Error>> {
1496 let actual = model.backend.model_identity(&model.model);
1497 if let Some(component) = model
1498 .backend
1499 .model_identity_mismatch(&self.model_identity, &actual)
1500 {
1501 return Err(RealtimeError::ModelMismatch { component });
1502 }
1503 Ok(())
1504 }
1505
1506 pub fn register_request(
1508 &mut self,
1509 model: &RealtimeModel<B>,
1510 request: RequestId,
1511 sampling: RealtimeSampling,
1512 ) -> Result<(), RealtimeError<B::Error>> {
1513 self.validate_model(model)?;
1514 self.scheduler.validate_registration(request)?;
1515 let state = model
1516 .backend
1517 .create_session(&model.model, sampling)
1518 .map_err(RealtimeError::Backend)?;
1519 self.scheduler.register(
1520 request,
1521 RealtimeSession {
1522 model_identity: self.model_identity.clone(),
1523 state,
1524 batch_size: None,
1525 },
1526 )?;
1527 Ok(())
1528 }
1529
1530 pub fn register_request_with_session(
1532 &mut self,
1533 model: &RealtimeModel<B>,
1534 request: RequestId,
1535 session: RealtimeSession<B>,
1536 ) -> Result<(), RealtimeError<B::Error>> {
1537 self.validate_model(model)?;
1538 self.scheduler.validate_registration(request)?;
1539 if let Some(component) = model
1540 .backend
1541 .model_identity_mismatch(&self.model_identity, &session.model_identity)
1542 {
1543 return Err(RealtimeError::ModelMismatch { component });
1544 }
1545 model
1546 .backend
1547 .validate_session(&model.model, &session.state)
1548 .map_err(RealtimeError::Backend)?;
1549 self.scheduler.register(request, session)?;
1550 Ok(())
1551 }
1552
1553 pub fn enqueue(
1555 &mut self,
1556 model: &RealtimeModel<B>,
1557 request: RequestId,
1558 input: B::Input,
1559 ) -> Result<WorkId, RealtimeError<B::Error>> {
1560 self.enqueue_with_deadline(model, request, input, None)
1561 }
1562
1563 pub fn enqueue_with_deadline(
1565 &mut self,
1566 model: &RealtimeModel<B>,
1567 request: RequestId,
1568 input: B::Input,
1569 deadline: Option<Instant>,
1570 ) -> Result<WorkId, RealtimeError<B::Error>> {
1571 self.validate_model(model)?;
1572 model
1573 .backend
1574 .validate_input(&model.model, &input)
1575 .map_err(RealtimeError::Backend)?;
1576 let batch = model.backend.input_batch_size(&input);
1577 self.validate_batch(request, batch)?;
1578 let work = self
1579 .scheduler
1580 .enqueue_with_deadline(request, input, deadline)?;
1581 self.scheduler
1582 .request_state_mut(request)?
1583 .batch_size
1584 .get_or_insert(batch);
1585 Ok(work)
1586 }
1587
1588 pub fn enqueue_batch(
1590 &mut self,
1591 model: &RealtimeModel<B>,
1592 request: RequestId,
1593 inputs: Vec<B::Input>,
1594 ) -> Result<Vec<WorkId>, RealtimeError<B::Error>> {
1595 self.validate_model(model)?;
1596 let mut expected = self
1597 .scheduler
1598 .request_state(request)
1599 .ok_or(SchedulerError::UnknownRequest(request))?
1600 .batch_size;
1601 for input in &inputs {
1602 model
1603 .backend
1604 .validate_input(&model.model, input)
1605 .map_err(RealtimeError::Backend)?;
1606 let actual = model.backend.input_batch_size(input);
1607 if let Some(expected) = expected {
1608 if actual != expected {
1609 return Err(RealtimeError::BatchSize {
1610 request: request.value(),
1611 expected,
1612 actual,
1613 });
1614 }
1615 } else {
1616 expected = Some(actual);
1617 }
1618 }
1619 let work = self.scheduler.enqueue_batch(request, inputs)?;
1620 if let Some(batch) = expected {
1621 self.scheduler
1622 .request_state_mut(request)?
1623 .batch_size
1624 .get_or_insert(batch);
1625 }
1626 Ok(work)
1627 }
1628
1629 fn validate_batch(
1630 &self,
1631 request: RequestId,
1632 actual: usize,
1633 ) -> Result<(), RealtimeError<B::Error>> {
1634 let state = self
1635 .scheduler
1636 .request_state(request)
1637 .ok_or(SchedulerError::UnknownRequest(request))?;
1638 if let Some(expected) = state.batch_size {
1639 if expected != actual {
1640 return Err(RealtimeError::BatchSize {
1641 request: request.value(),
1642 expected,
1643 actual,
1644 });
1645 }
1646 }
1647 Ok(())
1648 }
1649
1650 pub fn run_queued(
1652 &mut self,
1653 model: &mut RealtimeModel<B>,
1654 ) -> Result<Vec<RealtimeCompletedStep<B::Output>>, RealtimeError<B::Error>> {
1655 self.run_bounded(model, usize::MAX)
1656 }
1657
1658 pub fn run_bounded(
1660 &mut self,
1661 model: &mut RealtimeModel<B>,
1662 max_frames: usize,
1663 ) -> Result<Vec<RealtimeCompletedStep<B::Output>>, RealtimeError<B::Error>> {
1664 self.validate_model(model)?;
1665 if max_frames == 0 {
1666 return Err(RealtimeError::EmptyRunBound);
1667 }
1668 let now = Instant::now();
1669 let mut progress = self.scheduler.poll_completions(now);
1670 self.scheduler.prepare_bounded(max_frames, now)?;
1671 let backend_name = model.backend.name().to_owned();
1672 let backend = &model.backend;
1673 let backend_model = &mut model.model;
1674 progress.newly_submitted = self.scheduler.submit_prepared(
1675 now,
1676 |_, input, session| -> Result<RealtimeTransition<B>, B::Error> {
1677 let submission = backend.submit_step(backend_model, &mut session.state, input)?;
1678 let retained_resources = backend.retained_resources(&submission.completion);
1679 Ok(RealtimeTransition {
1680 backend_name: backend_name.clone(),
1681 retained_resources,
1682 output: submission.output,
1683 completion: submission.completion,
1684 })
1685 },
1686 )?;
1687 let completed = self.scheduler.poll_completions(now);
1688 progress.committed.extend(completed.committed);
1689 progress.failed.extend(completed.failed);
1690 if let Some((work, failure)) = progress.failed.first() {
1691 return Err(RealtimeError::Asynchronous {
1692 work: *work,
1693 message: failure.to_string(),
1694 });
1695 }
1696 Ok(progress
1697 .committed
1698 .into_iter()
1699 .map(|(work, _, transition)| RealtimeCompletedStep {
1700 work,
1701 output: transition.output,
1702 })
1703 .collect())
1704 }
1705
1706 pub fn finish_request(&mut self, request: RequestId) -> Result<(), RealtimeError<B::Error>> {
1708 self.scheduler.finish(request)?;
1709 Ok(())
1710 }
1711 pub fn cancel_request(&mut self, request: RequestId) -> Result<(), RealtimeError<B::Error>> {
1713 self.scheduler.cancel(request)?;
1714 Ok(())
1715 }
1716 pub fn release_request(
1718 &mut self,
1719 request: RequestId,
1720 ) -> Result<RealtimeSession<B>, RealtimeError<B::Error>> {
1721 Ok(self.scheduler.release(request)?)
1722 }
1723 pub fn forget_terminal_request(
1725 &mut self,
1726 request: RequestId,
1727 ) -> Result<RequestStatus, RealtimeError<B::Error>> {
1728 Ok(self.scheduler.forget_terminal(request)?)
1729 }
1730 pub fn request_status(&self, request: RequestId) -> Option<RequestStatus> {
1732 self.scheduler.request_status(request)
1733 }
1734 pub fn queued_for_request(&self, request: RequestId) -> usize {
1736 self.scheduler.queued_for_request(request)
1737 }
1738 pub fn set_request_sampling(
1740 &mut self,
1741 model: &RealtimeModel<B>,
1742 request: RequestId,
1743 sampling: RealtimeSampling,
1744 ) -> Result<(), RealtimeError<B::Error>> {
1745 self.validate_model(model)?;
1746 let queued = self.scheduler.queued_for_request(request);
1747 if queued != 0 {
1748 return Err(RealtimeError::SamplingWhileQueued {
1749 request: request.value(),
1750 queued,
1751 });
1752 }
1753 let state = self.scheduler.request_state_mut(request)?;
1754 model
1755 .backend
1756 .set_sampling(&mut state.state, sampling)
1757 .map_err(RealtimeError::Backend)
1758 }
1759 pub fn report(&self) -> SchedulerReport {
1761 self.scheduler.report()
1762 }
1763 pub fn capabilities(&self) -> SchedulerCapabilities {
1765 self.scheduler.capabilities()
1766 }
1767}
1768
1769#[cfg(test)]
1770mod tests {
1771 use super::*;
1772 use std::convert::Infallible;
1773
1774 #[derive(Clone)]
1775 struct MockSession {
1776 step: u32,
1777 sampling: RealtimeSampling,
1778 }
1779 impl SemanticStateTransaction for MockSession {
1780 type Branch = Self;
1781 type Error = Infallible;
1782 fn branch(&self) -> Result<Self, Self::Error> {
1783 Ok(self.clone())
1784 }
1785 fn commit_branch(&mut self, branch: Self) -> Result<(), Self::Error> {
1786 *self = branch;
1787 Ok(())
1788 }
1789 }
1790
1791 #[derive(Clone)]
1792 struct Frame(Vec<u32>);
1793 impl WorkDescriptor for Frame {
1794 type Error = Infallible;
1795 fn encode_descriptor(&self, output: &mut Vec<u32>) -> Result<(), Self::Error> {
1796 output.extend_from_slice(&self.0);
1797 Ok(())
1798 }
1799 }
1800
1801 struct Done;
1802 impl Completion for Done {
1803 type Error = Infallible;
1804 fn is_complete(&self) -> Result<bool, Self::Error> {
1805 Ok(true)
1806 }
1807 fn wait(&self) -> Result<(), Self::Error> {
1808 Ok(())
1809 }
1810 }
1811
1812 struct MockBackend;
1813 impl RealtimeBackend for MockBackend {
1814 type Model = u64;
1815 type ModelIdentity = u64;
1816 type Input = Frame;
1817 type Output = u32;
1818 type Session = MockSession;
1819 type Completion = Done;
1820 type Error = Infallible;
1821
1822 fn name(&self) -> &str {
1823 "mock-realtime"
1824 }
1825 fn model_identity(&self, model: &u64) -> u64 {
1826 *model
1827 }
1828 fn session_capabilities(&self, _: &u64) -> crate::SessionCapabilities {
1829 crate::SessionCapabilities::new(true, true, false)
1830 }
1831 fn speech_config(&self, _: &u64) -> RealtimeSpeechConfig {
1832 RealtimeSpeechConfig::new(
1833 2,
1834 1,
1835 1,
1836 1,
1837 0,
1838 0,
1839 RealtimeFrameConvention::FeedbackAlignedHistory,
1840 vec![0, 0, 1],
1841 )
1842 .unwrap()
1843 }
1844 fn materialize_input(
1845 &self,
1846 _: &u64,
1847 frame: &RealtimeInputFrame,
1848 ) -> Result<Frame, Infallible> {
1849 Ok(Frame(
1850 frame
1851 .input_audio_tokens()
1852 .iter()
1853 .map(|token| *token as u32)
1854 .collect(),
1855 ))
1856 }
1857 fn observe_output(&self, output: &u32) -> Result<RealtimeOutputFrame, Infallible> {
1858 Ok(RealtimeOutputFrame::new(
1859 1,
1860 vec![*output as i32],
1861 Vec::new(),
1862 Vec::new(),
1863 None,
1864 Vec::new(),
1865 ))
1866 }
1867 fn create_session(
1868 &self,
1869 _: &u64,
1870 sampling: RealtimeSampling,
1871 ) -> Result<MockSession, Infallible> {
1872 Ok(MockSession { step: 0, sampling })
1873 }
1874 fn validate_session(&self, _: &u64, _: &MockSession) -> Result<(), Infallible> {
1875 Ok(())
1876 }
1877 fn validate_input(&self, _: &u64, _: &Frame) -> Result<(), Infallible> {
1878 Ok(())
1879 }
1880 fn input_batch_size(&self, input: &Frame) -> usize {
1881 input.0.len()
1882 }
1883 fn set_sampling(
1884 &self,
1885 session: &mut MockSession,
1886 sampling: RealtimeSampling,
1887 ) -> Result<(), Infallible> {
1888 session.sampling = sampling;
1889 Ok(())
1890 }
1891 fn submit_step(
1892 &self,
1893 model: &mut u64,
1894 session: &mut MockSession,
1895 input: &Frame,
1896 ) -> Result<Submission<u32, Done>, Infallible> {
1897 session.step += 1;
1898 Ok(Submission {
1899 output: *model as u32 + session.step + input.0.iter().sum::<u32>(),
1900 completion: Done,
1901 })
1902 }
1903 }
1904
1905 impl RealtimeModelLoadingBackend for MockBackend {
1906 type Preparation = u64;
1907 type LoadOptions = u64;
1908
1909 fn materialize_realtime_model(
1910 &self,
1911 preparation: Self::Preparation,
1912 _: Self::LoadOptions,
1913 ) -> Result<Self::Model, Self::Error> {
1914 Ok(preparation)
1915 }
1916 }
1917
1918 #[test]
1919 fn selected_backend_materializes_architecture_preparation() {
1920 let model = load_realtime_model_with_options(MockBackend, 37, 0).unwrap();
1921 assert_eq!(*model.model(), 37);
1922 assert_eq!(model.backend().name(), "mock-realtime");
1923 assert_eq!(
1924 model.session_capabilities(),
1925 crate::SessionCapabilities::new(true, true, false)
1926 );
1927 }
1928
1929 #[test]
1930 fn mock_backend_runs_fair_realtime_sessions_without_accelerator_types() {
1931 let mut model = RealtimeModel::new(MockBackend, 10);
1932 let limits = SchedulerLimits::with_execution_bounds(2, 4, 2, 2, 1, usize::MAX).unwrap();
1933 let mut scheduler = RealtimeScheduler::new(&model, limits).unwrap();
1934 let first = RequestId::new(1);
1935 let second = RequestId::new(2);
1936 scheduler
1937 .register_request(&model, first, RealtimeSampling::greedy())
1938 .unwrap();
1939 scheduler
1940 .register_request(&model, second, RealtimeSampling::greedy())
1941 .unwrap();
1942 scheduler.enqueue(&model, first, Frame(vec![1])).unwrap();
1943 scheduler.enqueue(&model, second, Frame(vec![2])).unwrap();
1944 assert!(matches!(
1945 scheduler.set_request_sampling(
1946 &model,
1947 first,
1948 RealtimeSampling::new(0.5, 0.5, 7).unwrap()
1949 ),
1950 Err(RealtimeError::SamplingWhileQueued { .. })
1951 ));
1952 assert_eq!(
1953 scheduler
1954 .run_queued(&mut model)
1955 .unwrap()
1956 .into_iter()
1957 .map(|step| step.into_parts().1)
1958 .collect::<Vec<_>>(),
1959 vec![12, 13]
1960 );
1961 let updated = RealtimeSampling::new(0.5, 0.5, 7).unwrap();
1962 scheduler
1963 .set_request_sampling(&model, first, updated)
1964 .unwrap();
1965 assert_eq!(
1966 scheduler.release_request(first).unwrap().state().sampling,
1967 updated
1968 );
1969 }
1970
1971 #[test]
1972 fn portable_scheduler_lifecycle_rejects_mismatch_and_preserves_resumed_state() {
1973 let mut model = RealtimeModel::new(MockBackend, 10);
1974 let other_model = RealtimeModel::new(MockBackend, 11);
1975 let limits = SchedulerLimits::with_execution_bounds(2, 4, 2, 2, 1, usize::MAX).unwrap();
1976 let mut scheduler = RealtimeScheduler::new(&model, limits).unwrap();
1977
1978 let cancelled = RequestId::new(10);
1979 scheduler
1980 .register_request(&model, cancelled, RealtimeSampling::greedy())
1981 .unwrap();
1982 scheduler
1983 .enqueue(&model, cancelled, Frame(vec![1]))
1984 .unwrap();
1985 assert!(matches!(
1986 scheduler.enqueue(&model, cancelled, Frame(vec![1, 2])),
1987 Err(RealtimeError::BatchSize {
1988 request: 10,
1989 expected: 1,
1990 actual: 2,
1991 })
1992 ));
1993 assert_eq!(scheduler.queued_for_request(cancelled), 1);
1994 scheduler.cancel_request(cancelled).unwrap();
1995 assert_eq!(
1996 scheduler.request_status(cancelled),
1997 Some(RequestStatus::Cancelled)
1998 );
1999 assert_eq!(scheduler.queued_for_request(cancelled), 0);
2000
2001 assert!(matches!(
2002 scheduler.register_request(
2003 &other_model,
2004 RequestId::new(11),
2005 RealtimeSampling::greedy()
2006 ),
2007 Err(RealtimeError::ModelMismatch { .. })
2008 ));
2009
2010 let original = RequestId::new(20);
2011 scheduler
2012 .register_request(&model, original, RealtimeSampling::greedy())
2013 .unwrap();
2014 scheduler.enqueue(&model, original, Frame(vec![2])).unwrap();
2015 assert_eq!(scheduler.run_queued(&mut model).unwrap()[0].output(), &13);
2016 let released = scheduler.release_request(original).unwrap();
2017 assert_eq!(released.state().step, 1);
2018 assert_eq!(released.batch_size(), Some(1));
2019
2020 let resumed = RequestId::new(21);
2021 scheduler
2022 .register_request_with_session(&model, resumed, released)
2023 .unwrap();
2024 scheduler.enqueue(&model, resumed, Frame(vec![3])).unwrap();
2025 assert_eq!(scheduler.run_queued(&mut model).unwrap()[0].output(), &15);
2026 let released = scheduler.release_request(resumed).unwrap();
2027 assert_eq!(released.state().step, 2);
2028
2029 let mut other_scheduler = RealtimeScheduler::new(&other_model, limits).unwrap();
2030 assert!(matches!(
2031 other_scheduler.register_request_with_session(
2032 &other_model,
2033 RequestId::new(22),
2034 released
2035 ),
2036 Err(RealtimeError::ModelMismatch { .. })
2037 ));
2038 }
2039
2040 #[test]
2041 fn sampling_and_speech_config_validate_portably() {
2042 assert!(RealtimeSampling::new(f32::NAN, 0.0, 0).is_err());
2043 let config = RealtimeSpeechConfig::new(
2044 4,
2045 2,
2046 2,
2047 3,
2048 11,
2049 12,
2050 RealtimeFrameConvention::AbsoluteDelayedSlots,
2051 vec![2, 0, 1, 2, 3],
2052 )
2053 .unwrap();
2054 assert_eq!(config.max_audio_delay(), 3);
2055 assert_eq!(config.max_delay(), 3);
2056 assert_eq!(config.text_delay(), 2);
2057 assert_eq!(config.generated_audio_codebooks(), 2);
2058 assert_eq!(
2059 serde_json::from_str::<RealtimeSpeechConfig>(&serde_json::to_string(&config).unwrap())
2060 .unwrap(),
2061 config
2062 );
2063 let sampling = RealtimeSampling::new(0.7, 0.9, 42)
2064 .unwrap()
2065 .with_top_k(Some(25), Some(250))
2066 .unwrap();
2067 assert_eq!(sampling.text_top_k(), Some(25));
2068 assert_eq!(sampling.audio_top_k(), Some(250));
2069 assert!(RealtimeSampling::greedy()
2070 .with_top_k(Some(0), None)
2071 .is_err());
2072 assert_eq!(
2073 serde_json::from_str::<RealtimeSampling>(&serde_json::to_string(&sampling).unwrap())
2074 .unwrap(),
2075 sampling
2076 );
2077 let frame = RealtimeInputFrame::new(1, vec![1, 2])
2078 .with_forced_generated_audio(vec![3, 4])
2079 .with_forced_text(vec![5])
2080 .with_diagnostics();
2081 assert_eq!(frame.input_audio_tokens(), [1, 2]);
2082 assert_eq!(frame.forced_generated_audio_tokens(), Some(&[3, 4][..]));
2083 assert_eq!(frame.forced_text_tokens(), Some(&[5][..]));
2084 assert!(frame.retains_diagnostics());
2085 assert!(RealtimeSpeechConfig::new(
2086 4,
2087 1,
2088 1,
2089 1,
2090 0,
2091 0,
2092 RealtimeFrameConvention::FeedbackAlignedHistory,
2093 vec![0; 5],
2094 )
2095 .is_err());
2096 assert!(RealtimeSpeechConfig::new(
2097 1,
2098 0,
2099 1,
2100 1,
2101 -1,
2102 0,
2103 RealtimeFrameConvention::FeedbackAlignedHistory,
2104 vec![0; 2],
2105 )
2106 .is_err());
2107 assert!(RealtimeSpeechConfig::new(
2108 1,
2109 0,
2110 1,
2111 1,
2112 0,
2113 0,
2114 RealtimeFrameConvention::FeedbackAlignedHistory,
2115 vec![0, MAX_REALTIME_FRAME_DELAY + 1],
2116 )
2117 .is_err());
2118 }
2119
2120 fn released_schedule(
2121 convention: RealtimeFrameConvention,
2122 depth_audio_codebooks: usize,
2123 ) -> RealtimeSpeechConfig {
2124 RealtimeSpeechConfig::new(
2125 16,
2126 8,
2127 8,
2128 depth_audio_codebooks,
2129 32_000,
2130 2_048,
2131 convention,
2132 vec![0, 0, 1, 1, 1, 1, 1, 1, 1, 0, 1, 1, 1, 1, 1, 1, 1],
2133 )
2134 .unwrap()
2135 }
2136
2137 #[test]
2138 fn released_feedback_schedule_covers_warmup_forcing_and_output_alignment() {
2139 let config = released_schedule(RealtimeFrameConvention::FeedbackAlignedHistory, 8);
2140 let mut state = RealtimeFrameScheduleState::new(config.clone());
2141 let first = state
2142 .advance(&config, &RealtimeFrameForcing::none(&config))
2143 .unwrap();
2144 assert!(first.model_call_required());
2145 assert_eq!(first.input_placements().len(), 8);
2146 assert_eq!(first.temporal_inputs().len(), 17);
2147 assert!(first
2148 .temporal_inputs()
2149 .iter()
2150 .all(|source| matches!(source, RealtimeTemporalSource::Padding(_))));
2151 assert_eq!(first.targets().len(), 9);
2152 assert!(first.output().is_none());
2153
2154 let forcing = RealtimeFrameForcing::new(
2155 true,
2156 vec![true, false, false, false, false, false, false, false],
2157 );
2158 let second = state.advance(&config, &forcing).unwrap();
2159 assert_eq!(second.frontier(), 1);
2160 assert_eq!(second.next_frontier(), 2);
2161 assert_eq!(second.output().unwrap().len(), 8);
2162 assert_eq!(second.targets()[0].source(), RealtimeTargetSource::Forced);
2163 assert_eq!(second.targets()[1].source(), RealtimeTargetSource::Forced);
2164 assert_eq!(second.targets()[2].source(), RealtimeTargetSource::Sampled);
2165 assert!(matches!(
2166 second.temporal_inputs()[2],
2167 RealtimeTemporalSource::Padding(RealtimeFrameSlot::Audio(1))
2168 ));
2169 assert_eq!(
2170 second.output().unwrap()[0],
2171 RealtimeSlotCoordinate::new(0, RealtimeFrameSlot::Audio(0))
2172 );
2173 }
2174
2175 #[test]
2176 fn released_absolute_schedule_has_initialization_step_and_absolute_targets() {
2177 let config = released_schedule(RealtimeFrameConvention::AbsoluteDelayedSlots, 16);
2178 let mut state = RealtimeFrameScheduleState::new(config.clone());
2179 let initialization = state
2180 .advance(&config, &RealtimeFrameForcing::none(&config))
2181 .unwrap();
2182 assert!(!initialization.model_call_required());
2183 assert!(initialization.temporal_inputs().is_empty());
2184 assert!(initialization.targets().is_empty());
2185 assert_eq!(initialization.warmup_padding().len(), 17);
2186 assert!(initialization.output().is_none());
2187
2188 let forcing = RealtimeFrameForcing::new(
2189 true,
2190 vec![true, true, false, false, false, false, false, false],
2191 );
2192 let first_model = state.advance(&config, &forcing).unwrap();
2193 assert!(first_model.model_call_required());
2194 assert_eq!(first_model.temporal_inputs().len(), 17);
2195 assert!(first_model.temporal_inputs().iter().all(|source| matches!(
2196 source,
2197 RealtimeTemporalSource::Occupied {
2198 occupancy: RealtimeSlotOccupancy::Padding,
2199 ..
2200 }
2201 )));
2202 assert_eq!(first_model.targets().len(), 17);
2203 assert_eq!(
2204 first_model.targets()[0].source(),
2205 RealtimeTargetSource::Forced
2206 );
2207 assert_eq!(
2208 first_model.targets()[1].source(),
2209 RealtimeTargetSource::Forced
2210 );
2211 assert_eq!(
2212 first_model.targets()[2].source(),
2213 RealtimeTargetSource::Existing(RealtimeSlotOccupancy::Padding)
2214 );
2215 assert!(first_model.output().is_none());
2216
2217 let second_model = state
2218 .advance(&config, &RealtimeFrameForcing::none(&config))
2219 .unwrap();
2220 assert_eq!(second_model.output().unwrap().len(), 8);
2221 assert_eq!(
2222 second_model.output().unwrap()[0],
2223 RealtimeSlotCoordinate::new(1, RealtimeFrameSlot::Audio(0))
2224 );
2225 assert_eq!(
2226 second_model.output().unwrap()[1],
2227 RealtimeSlotCoordinate::new(2, RealtimeFrameSlot::Audio(1))
2228 );
2229 }
2230
2231 #[test]
2232 fn frame_schedule_branch_commit_and_rollback_are_atomic() {
2233 let config = released_schedule(RealtimeFrameConvention::FeedbackAlignedHistory, 8);
2234 let mut state = RealtimeFrameScheduleState::new(config.clone());
2235 let mut discarded = state.branch().unwrap();
2236 discarded
2237 .advance(&config, &RealtimeFrameForcing::none(&config))
2238 .unwrap();
2239 assert_eq!(state.frontier(), 0);
2240 RealtimeFrameScheduleState::discard_branch(discarded).unwrap();
2241
2242 let mut committed = state.branch().unwrap();
2243 committed
2244 .advance(&config, &RealtimeFrameForcing::none(&config))
2245 .unwrap();
2246 state.commit_branch(committed).unwrap();
2247 assert_eq!(state.frontier(), 1);
2248
2249 let other = released_schedule(RealtimeFrameConvention::AbsoluteDelayedSlots, 16);
2250 assert_eq!(
2251 state.validate_schedule(&other),
2252 Err(RealtimeScheduleError::ScheduleMismatch)
2253 );
2254 }
2255
2256 #[test]
2257 fn frame_schedule_rejects_masks_missing_history_and_coordinate_overflow_atomically() {
2258 let config = released_schedule(RealtimeFrameConvention::AbsoluteDelayedSlots, 16);
2259 let mut state = RealtimeFrameScheduleState::new(config.clone());
2260 assert!(matches!(
2261 state.advance(&config, &RealtimeFrameForcing::new(false, vec![false; 7])),
2262 Err(RealtimeScheduleError::ForcingCount { .. })
2263 ));
2264 assert_eq!(state.frontier(), 0);
2265
2266 let before = state.clone();
2267 state.frontier = usize::MAX;
2268 let overflow_before = state.clone();
2269 assert_eq!(
2270 state.advance(&config, &RealtimeFrameForcing::none(&config)),
2271 Err(RealtimeScheduleError::CoordinateOverflow)
2272 );
2273 assert_eq!(state, overflow_before);
2274
2275 let mut missing = before;
2276 missing.frontier = 1;
2277 let missing_before = missing.clone();
2278 assert!(matches!(
2279 missing.advance(&config, &RealtimeFrameForcing::none(&config)),
2280 Err(RealtimeScheduleError::MissingSlot { .. })
2281 ));
2282 assert_eq!(missing, missing_before);
2283 }
2284}