1use std::collections::BTreeMap;
4
5use eredu_core::{
6 InputExtent, InputMetadataKey, InputModality, InputPartDescriptor, InputPayloadKind,
7 InputTensorIdentity, PreparedInputError, PreparedInputIdentity,
8};
9
10#[derive(Debug, Clone, Eq, PartialEq)]
12#[non_exhaustive]
13pub enum PreparedInputPayload<Tensor> {
14 TokenIds(Tensor),
16 Tensor(Tensor),
18 Embeddings(Tensor),
20}
21
22impl<Tensor> PreparedInputPayload<Tensor> {
23 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 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#[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 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 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 pub const fn modality(&self) -> InputModality {
104 self.modality
105 }
106
107 pub const fn payload(&self) -> &PreparedInputPayload<Tensor> {
109 &self.payload
110 }
111
112 pub const fn metadata(&self) -> &BTreeMap<InputMetadataKey, Tensor> {
114 &self.metadata
115 }
116
117 pub fn metadata_value(&self, key: InputMetadataKey) -> Option<&Tensor> {
119 self.metadata.get(&key)
120 }
121
122 pub fn extents(&self) -> &[InputExtent] {
124 &self.extents
125 }
126
127 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#[derive(Debug, Clone, Eq, PartialEq)]
150pub struct PreparedModelInput<Tensor> {
151 parts: Vec<PreparedInputPart<Tensor>>,
152 identity: PreparedInputIdentity,
153}
154
155impl<Tensor> PreparedModelInput<Tensor> {
156 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 pub const fn identity(&self) -> &PreparedInputIdentity {
172 &self.identity
173 }
174
175 pub fn parts(&self) -> &[PreparedInputPart<Tensor>] {
177 &self.parts
178 }
179
180 pub fn len(&self) -> usize {
182 self.parts.len()
183 }
184
185 pub fn is_empty(&self) -> bool {
187 self.parts.is_empty()
188 }
189
190 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 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 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}