Skip to main content

eredu_runtime/
input.rs

1//! Backend-neutral ownership of prepared multimodal tensors.
2
3use std::collections::BTreeMap;
4
5use eredu_core::{
6    InputExtent, InputMetadataKey, InputModality, InputPartDescriptor, InputPayloadKind,
7    InputTensorIdentity, PreparedInputError, PreparedInputIdentity,
8};
9
10/// Primary tensor and its semantic role for one prepared input part.
11#[derive(Debug, Clone, Eq, PartialEq)]
12#[non_exhaustive]
13pub enum PreparedInputPayload<Tensor> {
14    /// Tokenizer vocabulary IDs.
15    TokenIds(Tensor),
16    /// Model-native features or patches that still require an encoder.
17    Tensor(Tensor),
18    /// Already projected decoder-width embeddings.
19    Embeddings(Tensor),
20}
21
22impl<Tensor> PreparedInputPayload<Tensor> {
23    /// Semantic payload kind.
24    pub const fn kind(&self) -> InputPayloadKind {
25        match self {
26            Self::TokenIds(_) => InputPayloadKind::TokenIds,
27            Self::Tensor(_) => InputPayloadKind::Tensor,
28            Self::Embeddings(_) => InputPayloadKind::Embeddings,
29        }
30    }
31
32    /// Borrows the backend-native tensor.
33    pub const fn value(&self) -> &Tensor {
34        match self {
35            Self::TokenIds(value) | Self::Tensor(value) | Self::Embeddings(value) => value,
36        }
37    }
38}
39
40/// One owned, typed, ordered prepared input part.
41#[derive(Debug, Clone, Eq, PartialEq)]
42pub struct PreparedInputPart<Tensor> {
43    modality: InputModality,
44    payload: PreparedInputPayload<Tensor>,
45    metadata: BTreeMap<InputMetadataKey, Tensor>,
46    extents: Vec<InputExtent>,
47}
48
49impl<Tensor> PreparedInputPart<Tensor> {
50    /// Creates a part with compatible payload and unique, compatible metadata.
51    pub fn new(
52        modality: InputModality,
53        payload: PreparedInputPayload<Tensor>,
54        metadata: impl IntoIterator<Item = (InputMetadataKey, Tensor)>,
55    ) -> Result<Self, PreparedInputError> {
56        Self::new_with_extents(modality, payload, metadata, [])
57    }
58
59    /// Creates a part with compatible host-known execution extents.
60    pub fn new_with_extents(
61        modality: InputModality,
62        payload: PreparedInputPayload<Tensor>,
63        metadata: impl IntoIterator<Item = (InputMetadataKey, Tensor)>,
64        extents: impl IntoIterator<Item = InputExtent>,
65    ) -> Result<Self, PreparedInputError> {
66        let payload_kind = payload.kind();
67        if !payload_kind.accepts(modality) {
68            return Err(PreparedInputError::IncompatiblePayload {
69                modality,
70                payload: payload_kind,
71            });
72        }
73        let mut typed_metadata = BTreeMap::new();
74        for (key, value) in metadata {
75            if !key.accepts(modality) {
76                return Err(PreparedInputError::IncompatibleMetadata { modality, key });
77            }
78            if typed_metadata.insert(key, value).is_some() {
79                return Err(PreparedInputError::DuplicateMetadata { key });
80            }
81        }
82        let extents = extents.into_iter().collect::<Vec<_>>();
83        for (index, extent) in extents.iter().copied().enumerate() {
84            if !extent.accepts(modality) {
85                return Err(PreparedInputError::IncompatibleExtent { modality, extent });
86            }
87            if extents[..index]
88                .iter()
89                .any(|prior| std::mem::discriminant(prior) == std::mem::discriminant(&extent))
90            {
91                return Err(PreparedInputError::DuplicateExtent { extent });
92            }
93        }
94        Ok(Self {
95            modality,
96            payload,
97            metadata: typed_metadata,
98            extents,
99        })
100    }
101
102    /// Part modality.
103    pub const fn modality(&self) -> InputModality {
104        self.modality
105    }
106
107    /// Primary tensor and semantic role.
108    pub const fn payload(&self) -> &PreparedInputPayload<Tensor> {
109        &self.payload
110    }
111
112    /// Typed metadata tensors in stable key order.
113    pub const fn metadata(&self) -> &BTreeMap<InputMetadataKey, Tensor> {
114        &self.metadata
115    }
116
117    /// Looks up one metadata tensor.
118    pub fn metadata_value(&self, key: InputMetadataKey) -> Option<&Tensor> {
119        self.metadata.get(&key)
120    }
121
122    /// Host-known extents needed by accelerator execution.
123    pub fn extents(&self) -> &[InputExtent] {
124        &self.extents
125    }
126
127    /// Builds and validates the core descriptor for this exact tensor part.
128    pub fn descriptor(
129        &self,
130        describe: &impl Fn(&Tensor) -> Result<InputTensorIdentity, PreparedInputError>,
131    ) -> Result<InputPartDescriptor, PreparedInputError> {
132        InputPartDescriptor::new_with_extents(
133            self.modality,
134            self.payload.kind(),
135            describe(self.payload.value())?,
136            self.metadata
137                .iter()
138                .map(|(key, value)| Ok((*key, describe(value)?)))
139                .collect::<Result<Vec<_>, PreparedInputError>>()?,
140            self.extents.iter().copied(),
141        )
142    }
143}
144
145/// Backend-neutral prepared input that owns backend-native tensor handles.
146///
147/// The identity is validated at construction and remains coupled to the exact
148/// ordered payload and metadata values used by runtime and distributed paths.
149#[derive(Debug, Clone, Eq, PartialEq)]
150pub struct PreparedModelInput<Tensor> {
151    parts: Vec<PreparedInputPart<Tensor>>,
152    identity: PreparedInputIdentity,
153}
154
155impl<Tensor> PreparedModelInput<Tensor> {
156    /// Validates and owns ordered prepared input parts.
157    pub fn new(
158        parts: Vec<PreparedInputPart<Tensor>>,
159        describe: impl Fn(&Tensor) -> Result<InputTensorIdentity, PreparedInputError>,
160    ) -> Result<Self, PreparedInputError> {
161        let identity = PreparedInputIdentity::new(
162            parts
163                .iter()
164                .map(|part| part.descriptor(&describe))
165                .collect::<Result<Vec<_>, _>>()?,
166        )?;
167        Ok(Self { parts, identity })
168    }
169
170    /// Exact payload-free identity used for rank agreement and persistence.
171    pub const fn identity(&self) -> &PreparedInputIdentity {
172        &self.identity
173    }
174
175    /// Ordered owned parts.
176    pub fn parts(&self) -> &[PreparedInputPart<Tensor>] {
177        &self.parts
178    }
179
180    /// Number of ordered parts.
181    pub fn len(&self) -> usize {
182        self.parts.len()
183    }
184
185    /// This input is always non-empty after construction.
186    pub fn is_empty(&self) -> bool {
187        self.parts.is_empty()
188    }
189
190    /// Borrows payloads and metadata tensors in deterministic wire order.
191    pub fn wire_values(&self) -> Vec<&Tensor> {
192        let mut values = Vec::new();
193        for part in &self.parts {
194            values.push(part.payload.value());
195            values.extend(part.metadata.values());
196        }
197        values
198    }
199
200    /// Reconstructs and validates input received in deterministic wire order.
201    pub fn from_identity_wire_values(
202        identity: PreparedInputIdentity,
203        values: Vec<Tensor>,
204        describe: impl Fn(&Tensor) -> Result<InputTensorIdentity, PreparedInputError>,
205    ) -> Result<Self, PreparedInputError> {
206        let expected_values = identity
207            .parts()
208            .iter()
209            .map(|part| 1 + part.metadata().len())
210            .sum::<usize>();
211        if values.len() != expected_values {
212            return Err(PreparedInputError::WireValueCount {
213                expected: expected_values,
214                actual: values.len(),
215            });
216        }
217        let mut values = values.into_iter();
218        let mut parts = Vec::with_capacity(identity.len());
219        for descriptor in identity.parts() {
220            let payload = values.next().expect("validated prepared-input value count");
221            let payload = match descriptor.payload_kind() {
222                InputPayloadKind::TokenIds => PreparedInputPayload::TokenIds(payload),
223                InputPayloadKind::Tensor => PreparedInputPayload::Tensor(payload),
224                InputPayloadKind::Embeddings => PreparedInputPayload::Embeddings(payload),
225                payload_kind => {
226                    return Err(PreparedInputError::IncompatiblePayload {
227                        modality: descriptor.modality(),
228                        payload: payload_kind,
229                    });
230                }
231            };
232            let metadata = descriptor
233                .metadata()
234                .keys()
235                .copied()
236                .map(|key| {
237                    (
238                        key,
239                        values.next().expect("validated prepared-input value count"),
240                    )
241                })
242                .collect::<Vec<_>>();
243            parts.push(PreparedInputPart::new_with_extents(
244                descriptor.modality(),
245                payload,
246                metadata,
247                descriptor.extents(),
248            )?);
249        }
250        let actual = Self::new(parts, describe)?;
251        if actual.identity != identity {
252            return Err(PreparedInputError::WireIdentityMismatch);
253        }
254        Ok(actual)
255    }
256
257    /// Consumes the lifecycle container and returns its ordered parts.
258    pub fn into_parts(self) -> Vec<PreparedInputPart<Tensor>> {
259        self.parts
260    }
261}
262
263#[cfg(test)]
264mod tests {
265    use eredu_core::{checkpoint::TensorDtype, PreparedInputError};
266
267    use super::*;
268
269    #[derive(Debug, Clone, Eq, PartialEq)]
270    struct FakeTensor {
271        dtype: TensorDtype,
272        shape: Vec<usize>,
273        marker: u8,
274    }
275
276    fn fake(dtype: TensorDtype, shape: &[usize], marker: u8) -> FakeTensor {
277        FakeTensor {
278            dtype,
279            shape: shape.to_vec(),
280            marker,
281        }
282    }
283
284    fn describe(value: &FakeTensor) -> Result<InputTensorIdentity, PreparedInputError> {
285        InputTensorIdentity::new(value.dtype.clone(), value.shape.clone())
286    }
287
288    #[test]
289    fn composite_input_extension_binds_typed_parts_to_a_multi_group_graph() {
290        let graph = crate::ExecutionGraph::new(
291            vec![
292                crate::ExecutionGroupSpec::root("vision"),
293                crate::ExecutionGroupSpec::with_dependencies("text", ["vision"]),
294            ],
295            "text",
296        )
297        .unwrap();
298        let input = PreparedModelInput::new(
299            vec![
300                PreparedInputPart::new(
301                    InputModality::Text,
302                    PreparedInputPayload::TokenIds(fake(TensorDtype::U32, &[1, 2], 1)),
303                    [],
304                )
305                .unwrap(),
306                PreparedInputPart::new_with_extents(
307                    InputModality::Image,
308                    PreparedInputPayload::Tensor(fake(TensorDtype::F32, &[4, 12], 2)),
309                    [(
310                        InputMetadataKey::PatchGrid,
311                        fake(TensorDtype::I32, &[1, 3], 3),
312                    )],
313                    [InputExtent::PatchGrid {
314                        time: 1,
315                        height: 2,
316                        width: 2,
317                    }],
318                )
319                .unwrap(),
320            ],
321            describe,
322        )
323        .unwrap();
324        let identity = input.identity().clone();
325        let values = input.wire_values().into_iter().cloned().collect();
326
327        let rebuilt =
328            PreparedModelInput::from_identity_wire_values(identity, values, describe).unwrap();
329        assert_eq!(rebuilt, input);
330        assert_eq!(graph.execution_order(), [0, 1]);
331        assert_eq!(graph.output(), 1);
332        assert_eq!(rebuilt.wire_values()[2].marker, 3);
333        assert_eq!(
334            rebuilt.parts()[1].extents(),
335            &[InputExtent::PatchGrid {
336                time: 1,
337                height: 2,
338                width: 2,
339            }]
340        );
341    }
342
343    #[test]
344    fn rejects_payload_geometry_that_disagrees_with_wire_identity() {
345        let input = PreparedModelInput::new(
346            vec![PreparedInputPart::new(
347                InputModality::Text,
348                PreparedInputPayload::TokenIds(fake(TensorDtype::U32, &[1, 2], 1)),
349                [],
350            )
351            .unwrap()],
352            describe,
353        )
354        .unwrap();
355        let wrong = vec![fake(TensorDtype::U32, &[1, 3], 1)];
356
357        assert!(matches!(
358            PreparedModelInput::from_identity_wire_values(
359                input.identity().clone(),
360                wrong,
361                describe
362            ),
363            Err(PreparedInputError::WireIdentityMismatch)
364        ));
365    }
366
367    #[test]
368    fn rejects_incompatible_payload_at_part_construction() {
369        let result = PreparedInputPart::new(
370            InputModality::Text,
371            PreparedInputPayload::Tensor(fake(TensorDtype::F32, &[1, 2], 1)),
372            [],
373        );
374
375        assert!(matches!(
376            result,
377            Err(PreparedInputError::IncompatiblePayload {
378                modality: InputModality::Text,
379                payload: InputPayloadKind::Tensor,
380            })
381        ));
382    }
383
384    #[test]
385    fn rejects_incompatible_metadata_at_part_construction() {
386        let result = PreparedInputPart::new(
387            InputModality::Text,
388            PreparedInputPayload::TokenIds(fake(TensorDtype::U32, &[1, 2], 1)),
389            [(
390                InputMetadataKey::PatchGrid,
391                fake(TensorDtype::I32, &[1, 3], 2),
392            )],
393        );
394
395        assert!(matches!(
396            result,
397            Err(PreparedInputError::IncompatibleMetadata {
398                modality: InputModality::Text,
399                key: InputMetadataKey::PatchGrid,
400            })
401        ));
402    }
403}