1use std::collections::BTreeSet;
4
5use eredu_core::InputModality;
6
7use crate::{
8 select_replicated_text_realization, BackendMechanismCapabilities, ReplicatedTextRequirements,
9 ReplicatedTextSelectionError, ReplicatedTextSelectionRequest,
10 SelectedReplicatedTextRealization,
11};
12
13#[derive(Debug, Clone, Copy, Eq, Hash, Ord, PartialEq, PartialOrd)]
15#[non_exhaustive]
16pub enum ProcessorPrimitive {
17 RgbResizeBicubic,
19 RgbResizeLanczos3,
21 RgbNormalize,
23 VideoSampling,
25 AudioWindow,
27 AudioSpectrum,
29 AudioMelFilter,
31 AudioLogarithm,
33 TensorU32,
35 TensorF32,
37 TensorI32,
39 TensorBool,
41 Padding,
43 Concatenation,
45 Slicing,
47 Indexing,
49 MaskConstruction,
51 Merge,
53 Encoder,
55 Projector,
57 MetadataInspection,
59}
60
61#[derive(Debug, Clone, Eq, PartialEq)]
63pub struct ModalityProcessorRequirements {
64 modality: InputModality,
65 raw_primitives: BTreeSet<ProcessorPrimitive>,
66 prepared_tensor: bool,
67 projected_embeddings: bool,
68 maximum_dimension: u64,
69}
70
71impl ModalityProcessorRequirements {
72 pub fn new(
74 modality: InputModality,
75 raw_primitives: impl IntoIterator<Item = ProcessorPrimitive>,
76 prepared_tensor: bool,
77 projected_embeddings: bool,
78 maximum_dimension: u64,
79 ) -> Result<Self, ProcessorSelectionError> {
80 if maximum_dimension == 0 {
81 return Err(ProcessorSelectionError::new([
82 "architecture declared a zero native dimension bound".into(),
83 ]));
84 }
85 Ok(Self {
86 modality,
87 raw_primitives: raw_primitives.into_iter().collect(),
88 prepared_tensor,
89 projected_embeddings,
90 maximum_dimension,
91 })
92 }
93
94 pub const fn modality(&self) -> InputModality {
96 self.modality
97 }
98
99 pub const fn raw_primitives(&self) -> &BTreeSet<ProcessorPrimitive> {
101 &self.raw_primitives
102 }
103
104 pub const fn prepared_tensor(&self) -> bool {
106 self.prepared_tensor
107 }
108
109 pub const fn projected_embeddings(&self) -> bool {
111 self.projected_embeddings
112 }
113
114 pub const fn maximum_dimension(&self) -> u64 {
116 self.maximum_dimension
117 }
118}
119
120#[derive(Debug, Clone, Eq, PartialEq)]
122pub struct ProcessorExecutionRequirements {
123 modalities: Vec<ModalityProcessorRequirements>,
124}
125
126impl ProcessorExecutionRequirements {
127 pub fn new(
129 modalities: impl IntoIterator<Item = ModalityProcessorRequirements>,
130 ) -> Result<Self, ProcessorSelectionError> {
131 let mut modalities = modalities.into_iter().collect::<Vec<_>>();
132 modalities.sort_by_key(|requirement| modality_order(requirement.modality));
133 if modalities
134 .windows(2)
135 .any(|pair| pair[0].modality == pair[1].modality)
136 {
137 return Err(ProcessorSelectionError::new([
138 "architecture repeats an input modality requirement".into(),
139 ]));
140 }
141 if modalities.is_empty() {
142 return Err(ProcessorSelectionError::new([
143 "architecture declares no input modalities".into(),
144 ]));
145 }
146 Ok(Self { modalities })
147 }
148
149 pub fn modalities(&self) -> &[ModalityProcessorRequirements] {
151 &self.modalities
152 }
153
154 pub fn modality(&self, modality: InputModality) -> Option<&ModalityProcessorRequirements> {
156 self.modalities
157 .iter()
158 .find(|requirement| requirement.modality == modality)
159 }
160}
161
162#[derive(Debug, Clone, Eq, PartialEq)]
164pub struct ProcessorSelectionRequest {
165 modalities: BTreeSet<InputModality>,
166 raw_media: bool,
167 available_raw_media: bool,
168 prepared_tensors: bool,
169 projected_modalities: BTreeSet<InputModality>,
170}
171
172impl ProcessorSelectionRequest {
173 pub fn new(modalities: impl IntoIterator<Item = InputModality>) -> Self {
175 Self {
176 modalities: modalities.into_iter().collect(),
177 raw_media: false,
178 available_raw_media: false,
179 prepared_tensors: true,
180 projected_modalities: BTreeSet::new(),
181 }
182 }
183
184 pub const fn with_raw_media(mut self, required: bool) -> Self {
186 self.raw_media = required;
187 self
188 }
189
190 pub const fn with_available_raw_media(mut self, enabled: bool) -> Self {
192 self.available_raw_media = enabled;
193 self
194 }
195
196 pub const fn with_prepared_tensors(mut self, required: bool) -> Self {
198 self.prepared_tensors = required;
199 self
200 }
201
202 pub fn with_projected_embeddings(mut self, required: bool) -> Self {
204 if required {
205 self.projected_modalities = self.modalities.clone();
206 } else {
207 self.projected_modalities.clear();
208 }
209 self
210 }
211
212 pub fn with_projected_modalities(
214 mut self,
215 modalities: impl IntoIterator<Item = InputModality>,
216 ) -> Self {
217 self.projected_modalities = modalities.into_iter().collect();
218 self
219 }
220
221 pub const fn modalities(&self) -> &BTreeSet<InputModality> {
223 &self.modalities
224 }
225}
226
227#[derive(Debug, Clone, Eq, PartialEq)]
229pub struct MediaPrimitiveCapabilities {
230 raw_modalities: BTreeSet<InputModality>,
231 prepared_modalities: BTreeSet<InputModality>,
232 projected_modalities: BTreeSet<InputModality>,
233 primitives: BTreeSet<ProcessorPrimitive>,
234 maximum_dimension: u64,
235}
236
237impl MediaPrimitiveCapabilities {
238 pub fn new(
240 raw_modalities: impl IntoIterator<Item = InputModality>,
241 prepared_modalities: impl IntoIterator<Item = InputModality>,
242 projected_modalities: impl IntoIterator<Item = InputModality>,
243 primitives: impl IntoIterator<Item = ProcessorPrimitive>,
244 maximum_dimension: u64,
245 ) -> Self {
246 Self {
247 raw_modalities: raw_modalities.into_iter().collect(),
248 prepared_modalities: prepared_modalities.into_iter().collect(),
249 projected_modalities: projected_modalities.into_iter().collect(),
250 primitives: primitives.into_iter().collect(),
251 maximum_dimension,
252 }
253 }
254
255 pub const fn primitives(&self) -> &BTreeSet<ProcessorPrimitive> {
257 &self.primitives
258 }
259}
260
261#[derive(Debug, Clone, Eq, PartialEq)]
263pub struct SelectedProcessorExecution {
264 requirements: ProcessorExecutionRequirements,
265 modalities: BTreeSet<InputModality>,
266 raw_media: bool,
267 prepared_tensors: bool,
268 projected_modalities: BTreeSet<InputModality>,
269}
270
271#[derive(Debug, Clone)]
273pub struct SelectedCompositeRealization {
274 execution: SelectedReplicatedTextRealization,
275 processor: SelectedProcessorExecution,
276}
277
278impl SelectedCompositeRealization {
279 pub fn from_parts(
281 execution: SelectedReplicatedTextRealization,
282 processor: SelectedProcessorExecution,
283 ) -> Self {
284 Self {
285 execution,
286 processor,
287 }
288 }
289
290 pub const fn execution(&self) -> &SelectedReplicatedTextRealization {
292 &self.execution
293 }
294
295 pub const fn processor(&self) -> &SelectedProcessorExecution {
297 &self.processor
298 }
299
300 pub fn into_parts(
302 self,
303 ) -> (
304 SelectedReplicatedTextRealization,
305 SelectedProcessorExecution,
306 ) {
307 (self.execution, self.processor)
308 }
309}
310
311#[derive(Debug, thiserror::Error)]
313pub enum CompositeSelectionError {
314 #[error(transparent)]
316 Execution(#[from] ReplicatedTextSelectionError),
317 #[error(transparent)]
319 Processor(#[from] ProcessorSelectionError),
320}
321
322pub fn select_composite_realization(
324 execution_requirements: &ReplicatedTextRequirements,
325 processor_requirements: &ProcessorExecutionRequirements,
326 execution_request: &ReplicatedTextSelectionRequest,
327 processor_request: &ProcessorSelectionRequest,
328 execution_capabilities: &BackendMechanismCapabilities,
329 processor_capabilities: &MediaPrimitiveCapabilities,
330) -> Result<SelectedCompositeRealization, CompositeSelectionError> {
331 let execution = select_replicated_text_realization(
332 execution_requirements,
333 execution_request,
334 execution_capabilities,
335 )?;
336 let processor = select_processor_execution(
337 processor_requirements,
338 processor_request,
339 processor_capabilities,
340 )?;
341 Ok(SelectedCompositeRealization::from_parts(
342 execution, processor,
343 ))
344}
345
346impl SelectedProcessorExecution {
347 pub const fn requirements(&self) -> &ProcessorExecutionRequirements {
349 &self.requirements
350 }
351
352 pub const fn modalities(&self) -> &BTreeSet<InputModality> {
354 &self.modalities
355 }
356
357 pub const fn raw_media(&self) -> bool {
359 self.raw_media
360 }
361
362 pub const fn prepared_tensors(&self) -> bool {
364 self.prepared_tensors
365 }
366
367 pub const fn projected_modalities(&self) -> &BTreeSet<InputModality> {
369 &self.projected_modalities
370 }
371}
372
373#[derive(Debug, Clone, Eq, PartialEq, thiserror::Error)]
375#[error("composite input realization is unsupported: {issues}", issues = .issues.join("; "))]
376pub struct ProcessorSelectionError {
377 issues: Vec<String>,
378}
379
380impl ProcessorSelectionError {
381 fn new(issues: impl IntoIterator<Item = String>) -> Self {
382 Self {
383 issues: issues.into_iter().collect(),
384 }
385 }
386
387 pub fn issues(&self) -> &[String] {
389 &self.issues
390 }
391}
392
393pub fn select_processor_execution(
395 requirements: &ProcessorExecutionRequirements,
396 request: &ProcessorSelectionRequest,
397 capabilities: &MediaPrimitiveCapabilities,
398) -> Result<SelectedProcessorExecution, ProcessorSelectionError> {
399 let mut issues = Vec::new();
400 let mut available_raw_media = request.available_raw_media
401 && request
402 .modalities
403 .iter()
404 .any(|modality| *modality != InputModality::Text);
405 for modality in &request.modalities {
406 let Some(requirement) = requirements.modality(*modality) else {
407 issues.push(format!("architecture input modality {}", modality.as_str()));
408 continue;
409 };
410 if requirement.maximum_dimension > capabilities.maximum_dimension {
411 issues.push(format!(
412 "{} native dimension {}",
413 modality.as_str(),
414 requirement.maximum_dimension
415 ));
416 }
417 if request.prepared_tensors
418 && (!requirement.prepared_tensor
419 || !capabilities.prepared_modalities.contains(modality))
420 {
421 issues.push(format!("{} prepared tensors", modality.as_str()));
422 }
423 if request.projected_modalities.contains(modality)
424 && (!requirement.projected_embeddings
425 || !capabilities.projected_modalities.contains(modality))
426 {
427 issues.push(format!("{} projected embeddings", modality.as_str()));
428 }
429 if (request.raw_media || request.available_raw_media) && *modality != InputModality::Text {
430 let mut raw_issues = Vec::new();
431 if requirement.raw_primitives.is_empty()
432 || !capabilities.raw_modalities.contains(modality)
433 {
434 raw_issues.push(format!("{} raw media", modality.as_str()));
435 }
436 for primitive in requirement
437 .raw_primitives
438 .difference(&capabilities.primitives)
439 {
440 raw_issues.push(format!("{} primitive {primitive:?}", modality.as_str()));
441 }
442 if raw_issues.is_empty() {
443 continue;
444 }
445 available_raw_media = false;
446 if request.raw_media {
447 issues.extend(raw_issues);
448 }
449 }
450 }
451 if issues.is_empty() {
452 Ok(SelectedProcessorExecution {
453 requirements: requirements.clone(),
454 modalities: request.modalities.clone(),
455 raw_media: request.raw_media || available_raw_media,
456 prepared_tensors: request.prepared_tensors,
457 projected_modalities: request.projected_modalities.clone(),
458 })
459 } else {
460 Err(ProcessorSelectionError::new(issues))
461 }
462}
463
464fn modality_order(modality: InputModality) -> u8 {
465 match modality {
466 InputModality::Text => 0,
467 InputModality::Image => 1,
468 InputModality::Video => 2,
469 InputModality::Audio => 3,
470 _ => u8::MAX,
471 }
472}
473
474#[cfg(test)]
475mod tests {
476 use super::*;
477
478 fn requirements() -> ProcessorExecutionRequirements {
479 ProcessorExecutionRequirements::new([
480 ModalityProcessorRequirements::new(
481 InputModality::Text,
482 [ProcessorPrimitive::TensorU32],
483 true,
484 true,
485 64,
486 )
487 .unwrap(),
488 ModalityProcessorRequirements::new(
489 InputModality::Image,
490 [
491 ProcessorPrimitive::RgbResizeBicubic,
492 ProcessorPrimitive::RgbNormalize,
493 ProcessorPrimitive::TensorF32,
494 ProcessorPrimitive::TensorI32,
495 ],
496 true,
497 true,
498 1024,
499 )
500 .unwrap(),
501 ])
502 .unwrap()
503 }
504
505 #[test]
506 fn selection_denies_each_missing_mechanism_without_callbacks() {
507 let request = ProcessorSelectionRequest::new([InputModality::Image])
508 .with_raw_media(true)
509 .with_projected_embeddings(true);
510 let capabilities = MediaPrimitiveCapabilities::new(
511 [InputModality::Image],
512 [InputModality::Image],
513 [InputModality::Image],
514 [
515 ProcessorPrimitive::RgbResizeBicubic,
516 ProcessorPrimitive::RgbNormalize,
517 ProcessorPrimitive::TensorF32,
518 ],
519 512,
520 );
521
522 let error = select_processor_execution(&requirements(), &request, &capabilities)
523 .expect_err("missing metadata construction and native extent must fail");
524 assert_eq!(
525 error.issues(),
526 ["image native dimension 1024", "image primitive TensorI32",]
527 );
528 }
529
530 #[test]
531 fn text_only_selection_does_not_require_unrequested_image_primitives() {
532 let request = ProcessorSelectionRequest::new([InputModality::Text]);
533 let capabilities = MediaPrimitiveCapabilities::new([], [InputModality::Text], [], [], 64);
534 let selected =
535 select_processor_execution(&requirements(), &request, &capabilities).unwrap();
536 assert_eq!(
537 selected.modalities(),
538 &BTreeSet::from([InputModality::Text])
539 );
540 assert!(!selected.raw_media());
541 }
542
543 #[test]
544 fn projected_readiness_requires_architecture_admission_and_backend_support() {
545 let requirements =
546 ProcessorExecutionRequirements::new([ModalityProcessorRequirements::new(
547 InputModality::Audio,
548 [],
549 true,
550 false,
551 64,
552 )
553 .unwrap()])
554 .unwrap();
555 let request =
556 ProcessorSelectionRequest::new([InputModality::Audio]).with_projected_embeddings(true);
557 let capabilities = MediaPrimitiveCapabilities::new(
558 [],
559 [InputModality::Audio],
560 [InputModality::Audio],
561 [],
562 64,
563 );
564 let error = select_processor_execution(&requirements, &request, &capabilities)
565 .expect_err("architecture-ineligible projected audio was admitted");
566 assert_eq!(error.issues(), ["audio projected embeddings"]);
567 }
568
569 #[test]
570 fn optional_raw_readiness_never_overstates_selected_mechanisms() {
571 let request = ProcessorSelectionRequest::new([InputModality::Image])
572 .with_available_raw_media(true)
573 .with_projected_embeddings(true);
574 let incomplete = MediaPrimitiveCapabilities::new(
575 [InputModality::Image],
576 [InputModality::Image],
577 [InputModality::Image],
578 [
579 ProcessorPrimitive::RgbResizeBicubic,
580 ProcessorPrimitive::RgbNormalize,
581 ProcessorPrimitive::TensorF32,
582 ],
583 1024,
584 );
585 let selected = select_processor_execution(&requirements(), &request, &incomplete).unwrap();
586 assert!(!selected.raw_media());
587 assert!(selected.prepared_tensors());
588 assert_eq!(
589 selected.projected_modalities(),
590 &BTreeSet::from([InputModality::Image])
591 );
592
593 let complete = MediaPrimitiveCapabilities::new(
594 [InputModality::Image],
595 [InputModality::Image],
596 [InputModality::Image],
597 [
598 ProcessorPrimitive::RgbResizeBicubic,
599 ProcessorPrimitive::RgbNormalize,
600 ProcessorPrimitive::TensorF32,
601 ProcessorPrimitive::TensorI32,
602 ],
603 1024,
604 );
605 let selected = select_processor_execution(&requirements(), &request, &complete).unwrap();
606 assert!(selected.raw_media());
607 assert!(selected.prepared_tensors());
608
609 let required = ProcessorSelectionRequest::new([InputModality::Image])
610 .with_raw_media(true)
611 .with_projected_embeddings(true);
612 let error = select_processor_execution(&requirements(), &required, &incomplete)
613 .expect_err("required raw readiness was admitted without every primitive");
614 assert_eq!(error.issues(), ["image primitive TensorI32"]);
615 }
616}