1use std::sync::Arc;
10
11use tenferro_ad::semantic_extension::{
12 AdValue, ResidualSpec, SemanticAdError, SemanticAdRuleRole, SemanticLinearizeRequest,
13 SemanticLinearizeResult, SemanticLinearizeRule, SemanticPrimalVjpRequest,
14 SemanticPrimalVjpRule,
15};
16use tenferro_ops::ext_op::ExtensionOp;
17use tenferro_runtime::program::SemanticProgramBuilder;
18
19use crate::extension::{
20 Df64Einsum, Df64EinsumJvp, Df64EinsumVjp, Df64Expand, Df64FromF64, Df64Qr, Df64QrJvp,
21 Df64QrVjp, Df64ToF64, Df64Total, DF64_OPS_FAMILY,
22};
23
24#[derive(Clone, Copy, Debug, PartialEq, Eq)]
26enum Df64Op {
27 Total,
29 Expand,
31 Qr,
33 Einsum,
35 ToF64,
37 FromF64,
39 QrVjp,
41 QrJvp,
43}
44
45impl Df64Op {
46 fn of(op: &dyn ExtensionOp) -> Option<Self> {
48 let any = op.as_any();
49 if any.downcast_ref::<Df64Total>().is_some() {
50 Some(Self::Total)
51 } else if any.downcast_ref::<Df64Expand>().is_some() {
52 Some(Self::Expand)
53 } else if any.downcast_ref::<Df64Einsum>().is_some() {
54 Some(Self::Einsum)
55 } else if any.downcast_ref::<Df64Qr>().is_some() {
56 Some(Self::Qr)
57 } else if any.downcast_ref::<Df64ToF64>().is_some() {
58 Some(Self::ToF64)
59 } else if any.downcast_ref::<Df64FromF64>().is_some() {
60 Some(Self::FromF64)
61 } else if any.downcast_ref::<Df64QrVjp>().is_some() {
62 Some(Self::QrVjp)
63 } else if any.downcast_ref::<Df64QrJvp>().is_some() {
64 Some(Self::QrJvp)
65 } else {
66 None
67 }
68 }
69}
70
71fn unsupported(op: Df64Op, role: SemanticAdRuleRole) -> SemanticAdError {
73 SemanticAdError::Rule {
74 family_id: DF64_OPS_FAMILY,
75 role,
76 source: Box::new(std::io::Error::other(format!(
77 "the Df64 operations family has no {role:?} rule for {op:?}"
78 ))),
79 }
80}
81
82#[derive(Debug)]
107pub struct Df64VjpRule;
108
109impl SemanticPrimalVjpRule for Df64VjpRule {
110 fn family_id(&self) -> &'static str {
111 DF64_OPS_FAMILY
112 }
113
114 fn residual_mask(&self) -> ResidualSpec {
115 ResidualSpec::all_outputs().with_all_inputs()
119 }
120
121 fn primal_vjp(
122 &self,
123 request: SemanticPrimalVjpRequest<'_>,
124 builder: &mut SemanticProgramBuilder,
125 ) -> Result<Box<[AdValue]>, SemanticAdError> {
126 let role = SemanticAdRuleRole::PrimalVjp;
127 let op = Df64Op::of(request.op()).ok_or_else(|| SemanticAdError::Rule {
128 family_id: DF64_OPS_FAMILY,
129 role,
130 source: Box::new(std::io::Error::other(
131 "the payload is not one of the Df64 operations",
132 )),
133 })?;
134 let inactive = || vec![AdValue::Absent; request.primal_input_count()].into_boxed_slice();
135 let emitted = match op {
136 Df64Op::Total => {
137 let Some(AdValue::Value(cotangent)) = request.cotangent_outputs().first().copied()
138 else {
139 return Ok(inactive());
140 };
141 let shape = exact_shape(&request)?;
145 builder.add_extension(Arc::new(Df64Expand::new(shape)), &[cotangent])
146 }
147 Df64Op::ToF64 | Df64Op::FromF64 => {
149 let Some(AdValue::Value(cotangent)) = request.cotangent_outputs().first().copied()
150 else {
151 return Ok(inactive());
152 };
153 let operation = if op == Df64Op::ToF64 {
154 Arc::new(Df64FromF64) as Arc<dyn ExtensionOp>
155 } else {
156 Arc::new(Df64ToF64) as Arc<dyn ExtensionOp>
157 };
158 builder.add_extension(operation, &[cotangent])
159 }
160 Df64Op::Qr => {
161 let has_q = matches!(request.cotangent_outputs().first(), Some(AdValue::Value(_)));
164 let has_r = matches!(request.cotangent_outputs().get(1), Some(AdValue::Value(_)));
165 if !has_q && !has_r {
166 return Ok(inactive());
167 }
168 let mut operands = vec![
169 request.primal_output_value(0)?,
170 request.primal_output_value(1)?,
171 ];
172 for (present, cotangent) in [
173 (has_q, request.cotangent_outputs().first().copied()),
174 (has_r, request.cotangent_outputs().get(1).copied()),
175 ] {
176 if !present {
177 continue;
178 }
179 match cotangent {
180 Some(AdValue::Value(value)) => operands.push(value),
181 _ => return Ok(inactive()),
182 }
183 }
184 builder.add_extension(Arc::new(Df64QrVjp::of(has_q, has_r)), &operands)
185 }
186 Df64Op::Einsum => {
187 let Some(contraction) = request.op().as_any().downcast_ref::<Df64Einsum>() else {
191 return Err(unsupported(op, role));
192 };
193 let Some(AdValue::Value(cotangent)) = request.cotangent_outputs().first().copied()
194 else {
195 return Ok(inactive());
196 };
197 let operands: Vec<&[u32]> = contraction
200 .input_labels()
201 .iter()
202 .map(|labels| labels.as_slice())
203 .collect();
204 let Ok(adjoint) = Df64EinsumVjp::of(&operands, contraction.out_labels()) else {
207 return Err(unsupported(op, role));
208 };
209 let mut call_operands = Vec::with_capacity(contraction.input_labels().len() + 1);
211 for index in 0..contraction.input_labels().len() {
212 call_operands.push(request.primal_input_value(index)?);
213 }
214 call_operands.push(cotangent);
215 builder.add_extension(Arc::new(adjoint), &call_operands)
216 }
217 Df64Op::Expand | Df64Op::QrVjp | Df64Op::QrJvp => {
218 return Err(unsupported(op, role));
219 }
220 }
221 .map_err(SemanticAdError::Build)?;
222
223 Ok(emitted
224 .iter()
225 .enumerate()
226 .map(|(index, value)| {
227 if index < request.primal_input_count() && request.active_inputs()[index] {
228 AdValue::Value(*value)
229 } else {
230 AdValue::Absent
231 }
232 })
233 .collect::<Vec<_>>()
234 .into_boxed_slice())
235 }
236}
237
238#[derive(Debug)]
255pub struct Df64LinearizeRule;
256
257impl SemanticLinearizeRule for Df64LinearizeRule {
258 fn family_id(&self) -> &'static str {
259 DF64_OPS_FAMILY
260 }
261
262 fn linearize(
263 &self,
264 request: SemanticLinearizeRequest<'_>,
265 builder: &mut SemanticProgramBuilder,
266 ) -> Result<SemanticLinearizeResult, SemanticAdError> {
267 let role = SemanticAdRuleRole::Linearize;
268 let op = Df64Op::of(request.op()).ok_or_else(|| SemanticAdError::Rule {
269 family_id: DF64_OPS_FAMILY,
270 role,
271 source: Box::new(std::io::Error::other(
272 "the payload is not one of the Df64 operations",
273 )),
274 })?;
275 let (operation, tangents): (Arc<dyn ExtensionOp>, Vec<_>) = match op {
276 Df64Op::Total => (
277 Arc::new(Df64Total),
278 request
279 .tangent_inputs()
280 .iter()
281 .filter_map(|value| value.value())
282 .collect(),
283 ),
284 Df64Op::ToF64 | Df64Op::FromF64 => {
285 let tangent = request
286 .tangent_inputs()
287 .iter()
288 .find_map(|value| value.value());
289 let operation: Arc<dyn ExtensionOp> = if op == Df64Op::ToF64 {
290 Arc::new(Df64ToF64)
291 } else {
292 Arc::new(Df64FromF64)
293 };
294 (operation, tangent.into_iter().collect())
295 }
296 Df64Op::Qr => {
297 let Some(tangent) = request
300 .tangent_inputs()
301 .first()
302 .and_then(|value| value.value())
303 else {
304 let inactive = (0..request.primal_outputs().len())
305 .map(|_| AdValue::Absent)
306 .collect::<Vec<_>>();
307 return Ok(SemanticLinearizeResult::new(inactive, Vec::new()));
308 };
309 (
310 Arc::new(Df64QrJvp) as Arc<dyn ExtensionOp>,
311 vec![
312 request.primal_outputs()[0],
313 request.primal_outputs()[1],
314 tangent,
315 ],
316 )
317 }
318 Df64Op::Einsum => {
319 let Some(contraction) = request.op().as_any().downcast_ref::<Df64Einsum>() else {
322 return Err(unsupported(op, role));
323 };
324 let has_lhs = request
325 .tangent_inputs()
326 .first()
327 .and_then(|value| value.value())
328 .is_some();
329 let has_rhs = request
330 .tangent_inputs()
331 .get(1)
332 .and_then(|value| value.value())
333 .is_some();
334 if !has_lhs && !has_rhs {
335 let inactive = (0..request.primal_inputs().len())
336 .map(|_| AdValue::Absent)
337 .collect::<Vec<_>>();
338 return Ok(SemanticLinearizeResult::new(inactive, Vec::new()));
339 }
340 let mask: Vec<bool> = request
344 .tangent_inputs()
345 .iter()
346 .map(|value| value.value().is_some())
347 .collect();
348 let operand_labels: Vec<&[u32]> = contraction
349 .input_labels()
350 .iter()
351 .map(|labels| labels.as_slice())
352 .collect();
353 let Ok(tangent) =
354 Df64EinsumJvp::of(&operand_labels, contraction.out_labels(), &mask)
355 else {
356 return Err(unsupported(op, role));
357 };
358 let mut operands: Vec<_> = request.primal_inputs().to_vec();
359 for value in request.tangent_inputs() {
360 if let Some(value) = value.value() {
361 operands.push(value);
362 }
363 }
364 (Arc::new(tangent) as Arc<dyn ExtensionOp>, operands)
365 }
366 Df64Op::Expand | Df64Op::QrVjp | Df64Op::QrJvp => {
367 return Err(unsupported(op, role));
368 }
369 };
370 if tangents.is_empty() {
371 let inactive = (0..request.primal_outputs().len())
372 .map(|_| AdValue::Absent)
373 .collect::<Vec<_>>();
374 return Ok(SemanticLinearizeResult::new(inactive, Vec::new()));
375 }
376 let emitted = builder
377 .add_extension(operation, &tangents)
378 .map_err(SemanticAdError::Build)?;
379 let outputs = emitted
380 .iter()
381 .enumerate()
382 .map(|(index, value)| {
383 if request
384 .active_outputs()
385 .get(index)
386 .copied()
387 .unwrap_or(false)
388 {
389 AdValue::Value(*value)
390 } else {
391 AdValue::Absent
392 }
393 })
394 .collect::<Vec<_>>();
395 Ok(SemanticLinearizeResult::new(outputs, Vec::new()))
396 }
397}
398
399fn exact_shape(request: &SemanticPrimalVjpRequest<'_>) -> Result<Vec<usize>, SemanticAdError> {
406 let metadata = request.primal_input_meta(0)?;
407 metadata
408 .shape()
409 .iter()
410 .map(|extent| match extent.as_exact() {
411 Some(tenferro_ops::dim_expr::DimExpr::Const(value)) => Ok(*value),
412 _ => Err(SemanticAdError::Invariant {
413 family_id: DF64_OPS_FAMILY,
414 role: SemanticAdRuleRole::PrimalVjp,
415 message: format!(
416 "the adjoint of {DF64_OPS_FAMILY} needs a concrete input shape, got {extent:?}"
417 ),
418 }),
419 })
420 .collect()
421}