1use anyhow::Result;
4use std::{collections::HashMap, rc::Rc};
5use syn::Ident;
6
7use proc_macro2::{Span, TokenStream};
8use quote::{format_ident, quote};
9
10use crate::ast_node::{LineHasSemicolon, Node};
11use crate::function_enums::BinaryMathFunctionEnum;
12use crate::operations::Operation;
13use crate::types::IJType;
14
15pub fn number_type_from_string(number_type: &str) -> Result<TokenStream, syn::Error> {
16 let parsed_val = syn::parse_str::<proc_macro2::TokenStream>(number_type).map_err(|err| {
17 syn::Error::new(
18 proc_macro2::Span::call_site(),
19 format!("Failed to parse value into a Rust TokenStream: {}", err),
20 )
21 })?;
22 Ok(parsed_val)
23}
24pub fn annotation_from_type(var_type: &IJType) -> Result<TokenStream> {
25 let number_type_annot = number_type_from_string(
26 &var_type
27 .extract_number_type()
28 .unwrap_or_default()
29 .unwrap_or("_".to_string()),
30 )?;
31 let res = match var_type {
32 IJType::Number(_) => quote!(#number_type_annot),
33 IJType::Tensor(_) => quote!(ijzer_lib::tensor::Tensor::<#number_type_annot>),
34 IJType::Function(signature) => {
35 let input_types = signature
36 .input
37 .iter()
38 .map(annotation_from_type)
39 .collect::<Result<Vec<_>>>()?;
40 let output_type = annotation_from_type(&signature.output)?;
41 quote! {
42 fn(#(#input_types),*) -> #output_type
43 }
44 }
45 IJType::Group(types) => {
46 let group_types = types
47 .iter()
48 .map(annotation_from_type)
49 .collect::<Result<Vec<_>>>()?;
50 quote!((#(#group_types),*))
51 }
52 _ => panic!("Unsupported type: {:?}", var_type),
53 };
54 Ok(res)
55}
56
57pub fn compile_line_from_node(
58 node: Rc<Node>,
59 has_semicolon: LineHasSemicolon,
60) -> Result<TokenStream> {
61 let mut compiler = CompilerContext::new(node.clone());
62 let mut line_stream: TokenStream = TokenStream::new();
63
64 while let Some(node_id) = compiler.pop_next_parseable() {
65 line_stream = compiler.compile(node_id)?;
66
67 compiler.submit_as_parsed(node_id, line_stream.clone())
68 }
69
70 if has_semicolon == LineHasSemicolon::Yes {
72 let mut ts_iter = line_stream.clone().into_iter().collect::<Vec<_>>();
73 let last_token = ts_iter.pop();
74 let needs_semicolon = !matches!(last_token, Some(proc_macro2::TokenTree::Punct(ref p)) if p.as_char() == ';' && p.spacing() == proc_macro2::Spacing::Alone)
75 && last_token.is_some();
76
77 if needs_semicolon {
78 line_stream = quote! {#line_stream;};
79 }
80 }
81
82 Ok(line_stream)
83}
84
85fn generate_identifiers(n: usize, name: &str) -> Vec<Ident> {
86 (1..=n).map(|i| format_ident!("{}{}", name, i)).collect()
87}
88
89pub struct CompilerContext {
90 pub node_map: HashMap<usize, Rc<Node>>,
91 parseable: Vec<usize>,
92 parsed: HashMap<usize, TokenStream>,
93 pub parent: HashMap<usize, usize>,
94 inputs: Vec<(usize, String)>,
95}
96
97impl CompilerContext {
98 pub fn new(root: Rc<Node>) -> Self {
99 let mut node_map = HashMap::new();
100 let mut leaves = vec![];
101 let mut parent = HashMap::new();
102 let mut stack = vec![root];
103
104 while let Some(node_rc) = stack.pop() {
105 node_map.insert(node_rc.id, node_rc.clone());
106 if node_rc.operands.is_empty() {
107 leaves.push(node_rc.id);
108 }
109 for child in &node_rc.operands {
110 parent.insert(child.id, node_rc.id);
111 stack.push(child.clone());
112 }
113 }
114 let parsed = HashMap::new();
115
116 Self {
117 node_map,
118 parseable: leaves,
119 parsed,
120 parent,
121 inputs: vec![],
122 }
123 }
124
125 pub fn submit_as_parsed(&mut self, node_id: usize, token_stream: TokenStream) {
126 self.parsed.insert(node_id, token_stream);
127
128 if let Some(&parent) = self.parent.get(&node_id) {
129 let siblings = self
130 .node_map
131 .get(&parent)
132 .unwrap()
133 .operands
134 .iter()
135 .map(|n| n.id)
136 .collect::<Vec<_>>();
137 if siblings.iter().all(|s| self.parsed.contains_key(s)) {
138 self.parseable.push(parent);
139 }
140 }
141 }
142
143 pub fn pop_next_parseable(&mut self) -> Option<usize> {
144 self.parseable.pop()
145 }
146
147 pub fn compile(&mut self, node_id: usize) -> Result<TokenStream> {
148 let node = self.node_map.get(&node_id).unwrap().clone();
149
150 let children = node.operands.iter().map(|n| n.id).collect::<Vec<_>>();
151 let child_streams: HashMap<usize, TokenStream> = children
152 .iter()
153 .map(|id| (*id, self.parsed.remove(id).unwrap()))
154 .collect();
155 let node_op = &node.op;
156
157 let stream = match node_op {
158 Operation::Number(_) => Number::compile(node, self, child_streams)?,
159 Operation::Symbol(_) => Symbol::compile(node, self, child_streams)?,
160 Operation::Function(_) => Function::compile(node, self, child_streams)?,
161 Operation::Group => Group::compile(node, self, child_streams)?,
162 Operation::Assign => Assign::compile(node, self, child_streams)?,
163 Operation::Identity => Identity::compile(node, self, child_streams)?,
164 Operation::Nothing => Nothing::compile(node, self, child_streams)?,
165 Operation::Subtract => Subtract::compile(node, self, child_streams)?,
166 Operation::Negate => Negate::compile(node, self, child_streams)?,
167 Operation::Array => Array::compile(node, self, child_streams)?,
168 Operation::Reduce => Reduce::compile(node, self, child_streams)?,
169 Operation::LambdaVariable(_) => LambdaVariable::compile(node, self, child_streams)?,
170 Operation::FunctionComposition(_) => {
171 FunctionComposition::compile(node, self, child_streams)?
172 }
173 Operation::Apply => Apply::compile(node, self, child_streams)?,
174 Operation::TypeConversion => TypeConversion::compile(node, self, child_streams)?,
175 Operation::GeneralizedContraction => {
176 GeneralizedContraction::compile(node, self, child_streams)?
177 }
178 Operation::TensorBuilder(_) => TensorBuilder::compile(node, self, child_streams)?,
179 Operation::Transpose => Transpose::compile(node, self, child_streams)?,
180 Operation::Shape => Shape::compile(node, self, child_streams)?,
181 Operation::QR => QR::compile(node, self, child_streams)?,
182 Operation::Svd => Svd::compile(node, self, child_streams)?,
183 Operation::Solve => Solve::compile(node, self, child_streams)?,
184 Operation::Diag => Diag::compile(node, self, child_streams)?,
185 Operation::Index => Index::compile(node, self, child_streams)?,
186 Operation::AssignSymbol(_) => AssignSymbol::compile(node, self, child_streams)?,
187 Operation::Reshape => Reshape::compile(node, self, child_streams)?,
188 Operation::UnaryFunction(_) => UnaryFunction::compile(node, self, child_streams)?,
189 Operation::BinaryFunction(_) => BinaryOperation::compile(node, self, child_streams)?,
190 Operation::Range => Range::compile(node, self, child_streams)?,
191 };
193
194 Ok(stream)
195 }
196
197 pub fn get_varname(&self, id: usize) -> Ident {
198 Ident::new(&format!("_{}", id), Span::call_site())
199 }
200}
201
202pub trait CompileNode {
203 fn compile(
204 node: Rc<Node>,
205 compiler: &mut CompilerContext,
206 child_streams: HashMap<usize, TokenStream>,
207 ) -> Result<TokenStream>;
208}
209
210struct Number;
211impl CompileNode for Number {
212 fn compile(
213 node: Rc<Node>,
214 _compiler: &mut CompilerContext,
215 _: HashMap<usize, TokenStream>,
216 ) -> Result<TokenStream> {
217 if let Operation::Number(number) = &node.op {
218 let val = number.value.clone();
219 let parsed_val = syn::parse_str::<proc_macro2::TokenStream>(&val).map_err(|_| {
220 syn::Error::new_spanned(&val, "Failed to parse value into a Rust TokenStream")
221 })?;
222 let number_type = node.output_type.extract_number_type().unwrap_or_default();
223 let res = match number_type {
224 Some(_) => {
225 let t = annotation_from_type(&IJType::Number(number_type))?;
226 quote! {
227 #parsed_val as #t
228 }
229 }
230 None => quote! {
231 #parsed_val
232 },
233 };
234 Ok(res)
235 } else {
236 panic!("Expected number node, found {:?}", node);
237 }
238 }
239}
240
241struct Array;
242impl CompileNode for Array {
243 fn compile(
244 node: Rc<Node>,
245 _compiler: &mut CompilerContext,
246 child_streams: HashMap<usize, TokenStream>,
247 ) -> Result<TokenStream> {
248 if let Operation::Array = &node.op {
249 } else {
250 panic!("Expected array node, found {:?}", node);
251 }
252 let child_ids = node.operands.iter().map(|n| n.id).collect::<Vec<_>>();
253 let child_streams = child_ids
254 .iter()
255 .map(|id| child_streams[id].clone())
256 .collect::<Vec<_>>();
257
258 let is_all_tensors = node
259 .operands
260 .iter()
261 .all(|n| n.output_type.type_match(&IJType::Tensor(None)));
262 let is_all_scalars = node
263 .operands
264 .iter()
265 .all(|n| n.output_type.type_match(&IJType::Number(None)));
266
267 let number_type = node.output_type.extract_number_type().unwrap_or_default();
268 let tensor_t = annotation_from_type(&IJType::Tensor(number_type))?;
269 if is_all_scalars {
270 let extracted_streams = child_streams
271 .into_iter()
272 .map(|s| {
273 quote! {#s}
274 })
275 .collect::<Vec<_>>();
276 let res = quote! {
277 #tensor_t::from_vec(vec![#(#extracted_streams),*], None)
278 };
279 Ok(res)
280 } else if is_all_tensors {
281 let res = quote! {
282 #tensor_t::from_tensors(&[#(#child_streams),*]).unwrap()
283 };
284 Ok(res)
285 } else {
286 panic!("Array elements must be all tensors or all scalars");
287 }
288 }
289}
290
291struct Symbol;
292impl CompileNode for Symbol {
293 fn compile(
294 node: Rc<Node>,
295 _: &mut CompilerContext,
296 _: HashMap<usize, TokenStream>,
297 ) -> Result<TokenStream> {
298 if let Operation::Symbol(symbol) = &node.op {
299 let varname = Ident::new(symbol.as_str(), Span::call_site());
300 match node.output_type {
301 IJType::Tensor(_) => Ok(quote! {
302 #varname.clone()
303 }),
304 _ => Ok(quote! {
305 #varname
306 }),
307 }
308 } else {
309 panic!("Expected symbol node, found {:?}", node);
310 }
311 }
312}
313struct AssignSymbol;
314impl CompileNode for AssignSymbol {
315 fn compile(
316 node: Rc<Node>,
317 _: &mut CompilerContext,
318 _: HashMap<usize, TokenStream>,
319 ) -> Result<TokenStream> {
320 if let Operation::AssignSymbol(symbol) = &node.op {
321 let varname = Ident::new(symbol.as_str(), Span::call_site());
322 Ok(quote! {
323 #varname
324 })
325 } else {
326 panic!("Expected symbol node, found {:?}", node);
327 }
328 }
329}
330
331struct Function;
332impl CompileNode for Function {
333 fn compile(
334 node: Rc<Node>,
335 _: &mut CompilerContext,
336 child_streams: HashMap<usize, TokenStream>,
337 ) -> Result<TokenStream> {
338 if let Operation::Function(function_name) = &node.op {
339 let varname = Ident::new(function_name.as_str(), Span::call_site());
340 let children = node.operands.iter().map(|n| n.id).collect::<Vec<_>>();
341
342 let child_streams: Vec<TokenStream> = children
343 .iter()
344 .map(|id| child_streams[id].clone())
345 .collect();
346
347 if child_streams.is_empty() {
348 Ok(quote! {#varname})
349 } else {
350 Ok(quote! {
351 #varname(#(#child_streams),*)
352 })
353 }
354 } else {
355 panic!("Expected symbol node, found {:?}", node);
356 }
357 }
358}
359
360struct Identity;
361impl CompileNode for Identity {
362 fn compile(
363 node: Rc<Node>,
364 _compiler: &mut CompilerContext,
365 _child_streams: HashMap<usize, TokenStream>,
366 ) -> Result<TokenStream> {
367 if let Operation::Identity = &node.op {
368 let children = node.operands.iter().map(|n| n.id).collect::<Vec<_>>();
369 if !children.is_empty() {
370 panic!("Expected 0 child, found {:?}", children.len());
371 }
372 _compiler.inputs.push((node.id, format!("_{}", node.id)));
373 let varname = _compiler.get_varname(node.id);
374
375 let res = quote! {
376 |#varname| #varname
377 };
378 Ok(res)
379 } else {
380 panic!("Expected identity node, found {:?}", node);
381 }
382 }
383}
384
385struct LambdaVariable;
386impl CompileNode for LambdaVariable {
387 fn compile(
388 node: Rc<Node>,
389 compiler: &mut CompilerContext,
390 _child_streams: HashMap<usize, TokenStream>,
391 ) -> Result<TokenStream> {
392 if let Operation::LambdaVariable(name) = &node.op {
393 let children = node.operands.iter().map(|n| n.id).collect::<Vec<_>>();
394 if !children.is_empty() {
395 panic!("Expected 0 child, found {:?}", children.len());
396 }
397 let name_with_underscores = format!("__{}", name);
398 compiler
399 .inputs
400 .push((node.id, name_with_underscores.clone()));
401 let varname = Ident::new(name_with_underscores.as_str(), Span::call_site());
402
403 let res = quote! {
404 #varname
405 };
406 Ok(res)
407 } else {
408 panic!("Expected lambda variable node, found {:?}", node);
409 }
410 }
411}
412
413struct Assign;
414impl CompileNode for Assign {
415 fn compile(
416 node: Rc<Node>,
417 _compiler: &mut CompilerContext,
418 child_streams: HashMap<usize, TokenStream>,
419 ) -> Result<TokenStream> {
420 if let Operation::Assign = &node.op {
421 } else {
422 panic!("Expected assign node, found {:?}", node);
423 }
424 let lhs_stream = child_streams[&node.operands[0].id].clone();
425 let lhs_annot = annotation_from_type(&node.operands[0].output_type)?;
426 let rhs_stream = child_streams[&node.operands[1].id].clone();
427 let args_streams = node
428 .operands
429 .iter()
430 .skip(2)
431 .map(|n| {
432 let s = child_streams[&n.id].clone();
433 let t = annotation_from_type(&n.output_type)?;
434 Ok(quote!(#s: #t))
435 })
436 .collect::<Result<Vec<_>>>()?;
437
438 if args_streams.is_empty() {
439 Ok(quote!(
440 let #lhs_stream: #lhs_annot = #rhs_stream;
441 ))
442 } else {
443 Ok(quote!(
444 let #lhs_stream = {|#(#args_streams),*| #rhs_stream};
445 ))
446 }
447 }
448}
449
450fn _compile_binary_op(
451 node: Rc<Node>,
452 child_streams: HashMap<usize, TokenStream>,
453 binary_op: impl Fn(TokenStream, TokenStream) -> TokenStream,
454) -> Result<TokenStream> {
455 let children = node.operands.iter().map(|n| n.id).collect::<Vec<_>>();
456
457 let res = match node.output_type.clone() {
458 IJType::Tensor(number_type) | IJType::Number(number_type) => {
459 if children.len() != 2 {
460 panic!("Expected 2 children for add, found {:?}", children.len());
461 }
462 let childstream1 = child_streams.get(&children[0]).unwrap();
463 let childstream2 = child_streams.get(&children[1]).unwrap();
464 let number_annot = annotation_from_type(&IJType::Number(number_type))?;
465 let operand1_type = node.operands[0].output_type.clone();
466 let operand2_type = node.operands[1].output_type.clone();
467
468 match (operand1_type, operand2_type) {
469 (IJType::Tensor(_), IJType::Tensor(_)) => {
470 let binop_stream = binary_op(quote! {_a}, quote! {_b});
471 quote! {
472 #childstream1.apply_binary_op(&#childstream2, |_a: #number_annot, _b: #number_annot| #binop_stream).unwrap()
473 }
474 }
475 (IJType::Tensor(_), IJType::Number(_)) => {
476 let binop_stream = binary_op(quote! {_x}, quote! {#childstream2});
477 quote! {
478 #childstream1.map(|_x: #number_annot| #binop_stream)
479 }
480 }
481 (IJType::Number(_), IJType::Tensor(_)) => {
482 let binop_stream = binary_op(quote! {#childstream1}, quote! {_x});
483 quote! {
484 #childstream2.map(|_x: #number_annot| #binop_stream)
485 }
486 }
487 (IJType::Number(_), IJType::Number(_)) => {
488 binary_op(quote! {#childstream1}, quote! {#childstream2})
489 }
490 _ => unreachable!(),
491 }
492 }
493 IJType::Function(f) => {
494 let number_type = f.output.extract_number_type().unwrap_or_default();
495 let number_annot = annotation_from_type(&IJType::Number(number_type.clone()))?;
496 let tensor_annot = annotation_from_type(&IJType::Tensor(number_type))?;
497 let input_types = f.input.clone();
498 if input_types.len() != 2 {
499 panic!(
500 "Expected 2 input types for add, found {:?}",
501 input_types.len()
502 );
503 }
504 let (input1_type, input2_type) = (input_types[0].clone(), input_types[1].clone());
505 match (input1_type, input2_type) {
506 (IJType::Number(_), IJType::Number(_)) => {
507 let binop_stream = binary_op(quote! {a}, quote! {b});
508 quote! { |a: #number_annot, b: #number_annot| #binop_stream }
509 }
510 (IJType::Tensor(_), IJType::Tensor(_)) => {
511 let binop_stream = binary_op(quote! {a}, quote! {b});
512 quote! {
513 |x1: #tensor_annot, x2: #tensor_annot|
514 x1.apply_binary_op(&x2, |a: #number_annot, b: #number_annot| #binop_stream).unwrap()
515 }
516 }
517 (IJType::Tensor(_), IJType::Number(_)) => {
518 let binop_stream = binary_op(quote! {a}, quote! {y});
519 quote! {
520 |x: #tensor_annot, y: #number_annot|
521 x.map(|a: #number_annot| #binop_stream)
522 }
523 }
524 (IJType::Number(_), IJType::Tensor(_)) => {
525 let binop_stream = binary_op(quote! {y}, quote! {a});
526 quote! {
527 |y: #number_annot, x: #tensor_annot|
528 x.map(|a: #number_annot| #binop_stream)
529 }
530 }
531 _ => panic!(
532 "Found add node with unimplemented input types: {:?}",
533 input_types
534 ),
535 }
536 }
537 _ => panic!(
538 "Found add node with unimplemented output type: {:?}",
539 node.output_type
540 ),
541 };
542 Ok(res)
543}
544
545struct BinaryOperation;
546impl CompileNode for BinaryOperation {
547 fn compile(
548 node: Rc<Node>,
549 _compiler: &mut CompilerContext,
550 child_streams: HashMap<usize, TokenStream>,
551 ) -> Result<TokenStream> {
552 let function = match node.op.clone() {
553 Operation::BinaryFunction(function) => function,
554 _ => unreachable!("Expected binary operation node, found {:?}", node.op),
555 };
556 let binary_op = match function {
557 BinaryMathFunctionEnum::Add => |a: TokenStream, b: TokenStream| quote! {#a + #b},
558 BinaryMathFunctionEnum::Multiply => |a: TokenStream, b: TokenStream| quote! {#a * #b},
559 BinaryMathFunctionEnum::Div => |a: TokenStream, b: TokenStream| quote! {#a / #b},
560 BinaryMathFunctionEnum::Power => |a: TokenStream, b: TokenStream| quote! {#a.pow(#b)},
561 BinaryMathFunctionEnum::Max => |a: TokenStream, b: TokenStream| quote! {#a.max(#b)},
562 BinaryMathFunctionEnum::Min => |a: TokenStream, b: TokenStream| quote! {#a.min(#b)},
563 BinaryMathFunctionEnum::Equals => {
564 |a: TokenStream, b: TokenStream| quote! {ijzer_lib::comparison_funcs::equals(#a, #b)}
565 }
566 BinaryMathFunctionEnum::NotEquals => {
567 |a: TokenStream, b: TokenStream| quote! {ijzer_lib::comparison_funcs::not_equals(#a, #b)}
568 }
569 BinaryMathFunctionEnum::GreaterThan => {
570 |a: TokenStream, b: TokenStream| quote! {ijzer_lib::comparison_funcs::greater_than(#a, #b)}
571 }
572 BinaryMathFunctionEnum::LessThan => {
573 |a: TokenStream, b: TokenStream| quote! {ijzer_lib::comparison_funcs::less_than(#a, #b)}
574 }
575 BinaryMathFunctionEnum::GreaterThanOrEqual => {
576 |a: TokenStream, b: TokenStream| quote! {ijzer_lib::comparison_funcs::greater_than_or_equal(#a, #b)}
577 }
578 BinaryMathFunctionEnum::LessThanOrEqual => {
579 |a: TokenStream, b: TokenStream| quote! {ijzer_lib::comparison_funcs::less_than_or_equal(#a, #b)}
580 }
581 BinaryMathFunctionEnum::Or => {
582 |a: TokenStream, b: TokenStream| quote! {ijzer_lib::comparison_funcs::or(#a, #b)}
583 }
584 BinaryMathFunctionEnum::And => {
585 |a: TokenStream, b: TokenStream| quote! {ijzer_lib::comparison_funcs::and(#a, #b)}
586 }
587 };
588 _compile_binary_op(node, child_streams, binary_op)
589 }
590}
591struct Subtract;
592impl CompileNode for Subtract {
593 fn compile(
594 node: Rc<Node>,
595 _compiler: &mut CompilerContext,
596 child_streams: HashMap<usize, TokenStream>,
597 ) -> Result<TokenStream> {
598 let binary_op = |a: TokenStream, b: TokenStream| quote! {#a - #b};
599 _compile_binary_op(node, child_streams, binary_op)
600 }
601}
602struct Negate;
603impl CompileNode for Negate {
604 fn compile(
605 node: Rc<Node>,
606 _compiler: &mut CompilerContext,
607 child_streams: HashMap<usize, TokenStream>,
608 ) -> Result<TokenStream> {
609 if let Operation::Negate = &node.op {
610 let children = node.operands.iter().map(|n| n.id).collect::<Vec<_>>();
611 let res = match node.output_type.clone() {
612 IJType::Tensor(number_type) => {
613 let number_annot = annotation_from_type(&IJType::Number(number_type))?;
614 let childstream = child_streams.get(&children[0]).unwrap();
615 quote! {
616 #childstream.map(|a: #number_annot| -a)
617 }
618 }
619 IJType::Number(_) => {
620 let childstream = child_streams.get(&children[0]).unwrap();
621 quote! {
622 -#childstream
623 }
624 }
625 IJType::Function(f) => {
626 let output_type = *f.output;
627 let number_type = output_type.extract_number_type().unwrap_or_default();
628 let number_annot = annotation_from_type(&IJType::Number(number_type.clone()))?;
629 let tensor_annot = annotation_from_type(&IJType::Tensor(number_type))?;
630 match output_type {
631 IJType::Number(_) => {
632 quote! { |a: #number_annot| -a }
633 }
634 IJType::Tensor(_) => {
635 quote! {
636 |x: #tensor_annot| x.map(|a: #number_annot| -a)
637 }
638 }
639 _ => panic!(
640 "Found negate node with unimplemented output type: {:?}",
641 node.output_type
642 ),
643 }
644 }
645 _ => panic!(
646 "Found negate node with unimplemented output type: {:?}",
647 node.output_type
648 ),
649 };
650 Ok(res)
651 } else {
652 panic!("Expected negate node, found {:?}", node);
653 }
654 }
655}
656
657struct Group;
658impl CompileNode for Group {
659 fn compile(
660 node: Rc<Node>,
661 _compiler: &mut CompilerContext,
662 child_streams: HashMap<usize, TokenStream>,
663 ) -> Result<TokenStream> {
664 if let Operation::Group = &node.op {
665 let children = node.operands.iter().map(|n| n.id).collect::<Vec<_>>();
666 let child_streams: Vec<TokenStream> = children
667 .iter()
668 .map(|id| child_streams[id].clone())
669 .collect();
670
671 match child_streams.len() {
672 1 => {
673 let stream = child_streams.first().unwrap();
674 Ok(quote! {#stream})
675 }
676 _ => Ok(quote! {
677 (#(#child_streams),*)
678 }),
679 }
680 } else {
681 panic!("Expected group node, found {:?}", node);
682 }
683 }
684}
685
686struct Reduce;
687impl CompileNode for Reduce {
688 fn compile(
689 node: Rc<Node>,
690 compiler: &mut CompilerContext,
691 child_streams: HashMap<usize, TokenStream>,
692 ) -> Result<TokenStream> {
693 if let Operation::Reduce = &node.op {
694 } else {
695 panic!("Expected reduce node, found {:?}", node.op);
696 }
697 let operand_ids = node.operands.iter().map(|n| n.id).collect::<Vec<_>>();
698 match operand_ids.len() {
699 1 => {
700 let functional_operand = &operand_ids[0];
701 let number_type = node
702 .output_type
703 .extract_signature()
704 .unwrap()
705 .output
706 .extract_number_type()
707 .unwrap_or_default();
708 let tensor_t = annotation_from_type(&IJType::Tensor(number_type))?;
709 let functional_operand_stream = child_streams.get(functional_operand).unwrap();
710 let ident = compiler.get_varname(node.id);
711 Ok(quote! {
712 |#ident: #tensor_t| #ident.reduce(#functional_operand_stream)
713 })
714 }
715 2 => {
716 let functional_operand = &operand_ids[0];
717 let functional_operand_stream = child_streams.get(functional_operand).unwrap();
718
719 let data_stream = child_streams.get(&operand_ids[1]).unwrap();
720
721 Ok(quote! {
722 #data_stream.reduce(#functional_operand_stream)
723 })
724 }
725 _ => {
726 panic!(
727 "Expected 1 or 2 operands for reduce operation, found {}",
728 operand_ids.len()
729 );
730 }
731 }
732 }
733}
734
735struct FunctionComposition;
736impl FunctionComposition {
737 fn apply_stream_to_stream(stream1: TokenStream, stream2: TokenStream) -> TokenStream {
738 quote! {
739 (#stream1)(#stream2)
740 }
741 }
742}
743impl CompileNode for FunctionComposition {
744 fn compile(
745 node: Rc<Node>,
746 _compiler: &mut CompilerContext,
747 child_streams: HashMap<usize, TokenStream>,
748 ) -> Result<TokenStream> {
749 let num_functions = if let Operation::FunctionComposition(n) = &node.op {
750 *n
751 } else {
752 panic!(
753 "Expected FunctionComposition operation, found {:?}",
754 node.op
755 );
756 };
757
758 let operand_ids = node.operands.iter().map(|n| n.id).collect::<Vec<_>>();
760 let functional_operand_ids = operand_ids.iter().take(num_functions);
761 let data_operand_ids = operand_ids.iter().skip(num_functions);
762 let data_streams = data_operand_ids
763 .map(|id| child_streams[id].clone())
764 .collect::<Vec<_>>();
765
766 let last_functional_operand = node.operands[num_functions - 1].clone();
767 let mut args = Vec::new();
768 let mut args_with_type = Vec::new();
769 for (i, arg_type) in last_functional_operand
770 .output_type
771 .extract_signature()
772 .unwrap()
773 .input
774 .iter()
775 .enumerate()
776 {
777 let ident = Ident::new(&format!("_{}_{}", node.id, i + 1), Span::call_site());
778 let type_annotation = annotation_from_type(arg_type)?;
779 args_with_type.push(quote! { #ident: #type_annotation });
780 args.push(quote! { #ident });
781 }
782 let closure = functional_operand_ids.rev().fold(
783 quote! {
784 #(#args),*
785 },
786 |acc, id| FunctionComposition::apply_stream_to_stream(child_streams[id].clone(), acc),
787 );
788 let closure = quote! {|#(#args_with_type),*| #closure};
789
790 if data_streams.is_empty() {
791 Ok(closure)
792 } else {
793 Ok(quote! {
794 (#closure)(#(#data_streams),*)
795 })
796 }
797 }
798}
799
800struct Nothing;
801impl CompileNode for Nothing {
802 fn compile(
803 node: Rc<Node>,
804 _compiler: &mut CompilerContext,
805 _: HashMap<usize, TokenStream>,
806 ) -> Result<TokenStream> {
807 if let Operation::Nothing = &node.op {
808 let res = quote! {};
809 Ok(res)
810 } else {
811 panic!("Expected nothing node, found {:?}", node);
812 }
813 }
814}
815
816#[allow(dead_code)]
817struct NotImplemented;
818impl CompileNode for NotImplemented {
819 fn compile(
820 node: Rc<Node>,
821 _compiler: &mut CompilerContext,
822 _: HashMap<usize, TokenStream>,
823 ) -> Result<TokenStream> {
824 let error_msg = format!("Compilation for operation '{:?}' not implemented", node.op);
825 Ok(quote! {
826 panic!(#error_msg);
827 })
828 }
829}
830
831struct Apply;
833impl CompileNode for Apply {
834 fn compile(
835 node: Rc<Node>,
836 _compiler: &mut CompilerContext,
837 child_streams: HashMap<usize, TokenStream>,
838 ) -> Result<TokenStream> {
839 if let Operation::Apply = &node.op {
840 } else {
841 panic!("Expected apply node, found {:?}", node.op);
842 }
843
844 let children = node.operands.iter().map(|n| n.id).collect::<Vec<_>>();
845 let child_streams: Vec<TokenStream> = children
846 .iter()
847 .map(|id| child_streams[id].clone())
848 .collect();
849
850 if child_streams.is_empty() {
851 panic!("Expected at least one child for apply node, found 0");
852 }
853
854 let function_stream = child_streams.first().unwrap();
855 let operand_streams = child_streams.iter().skip(1);
856
857 Ok(quote! {
858 (#function_stream)(#(#operand_streams),*)
859 })
860 }
861}
862
863struct TypeConversion;
864impl TypeConversion {
865 fn convert_type(
866 from: &IJType,
867 to: &IJType,
868 child_stream: TokenStream,
869 id: usize,
870 ) -> Result<TokenStream> {
871 let res = match (from, to) {
872 (IJType::Tensor(_), IJType::Tensor(_)) => quote! {#child_stream},
873 (IJType::Number(_), IJType::Number(_)) => quote! {#child_stream},
874 (IJType::Number(_), IJType::Tensor(number_type)) => {
875 let tensor_t = annotation_from_type(&IJType::Tensor(number_type.clone()))?;
876 quote! {#tensor_t::scalar(#child_stream)}
877 }
878 (IJType::Function(ref signature_from), IJType::Function(ref signature_to)) => {
879 let num_ops = signature_from.input.len();
880 let idents = generate_identifiers(num_ops, &format!("_{}_", id));
881
882 let mut input_conversions = Vec::new();
883 let mut args = Vec::new();
884 for ((from, to), ident) in signature_from
885 .input
886 .iter()
887 .zip(signature_to.input.clone())
888 .zip(idents.clone())
889 {
890 input_conversions.push(Self::convert_type(&to, from, quote!(#ident), id)?);
891 let arg_type = annotation_from_type(&to)?;
892 args.push(quote!(#ident: #arg_type));
893 }
894 let input_stream = quote! {(#child_stream)(#(#input_conversions),*)};
895 let converted_stream = Self::convert_type(
896 &signature_from.output,
897 &signature_to.output,
898 input_stream,
899 id,
900 )?;
901 quote! {
902 (|#(#args),*| (#converted_stream))
903 }
904 }
905 _ => {
906 panic!("Type conversion from {} to {} not implemented", from, to);
907 }
908 };
909 Ok(res)
910 }
911}
912impl CompileNode for TypeConversion {
913 fn compile(
914 node: Rc<Node>,
915 _compiler: &mut CompilerContext,
916 child_streams: HashMap<usize, TokenStream>,
917 ) -> Result<TokenStream> {
918 if node.operands.len() != 1 {
919 panic!("Expected 1 operands, found {}", node.operands.len());
920 }
921 let child_stream = child_streams.get(&node.operands[0].id).unwrap().clone();
922 Self::convert_type(
923 &node.input_types[0],
924 &node.output_type,
925 child_stream,
926 node.id,
927 )
928 }
929}
930
931struct GeneralizedContraction;
932impl CompileNode for GeneralizedContraction {
933 fn compile(
934 node: Rc<Node>,
935 _compiler: &mut CompilerContext,
936 child_streams: HashMap<usize, TokenStream>,
937 ) -> Result<TokenStream> {
938 let child_streams = node
939 .operands
940 .iter()
941 .map(|n| child_streams[&n.id].clone())
942 .collect::<Vec<_>>();
943 let f_stream = child_streams[0].clone();
944 let f_node = &node.operands[0];
945 let f_arg_annot =
946 annotation_from_type(&f_node.output_type.extract_signature().unwrap().input[0])?;
947 let f_stream_extract = quote! {
948 (|z: &#f_arg_annot| (#f_stream)(z.clone()))
949 };
950 let g_stream = child_streams[1].clone();
951 let g_node = &node.operands[1];
952 let input_number_type = g_node.output_type.extract_signature().unwrap().input[0]
953 .extract_number_type()
954 .unwrap_or_default();
955 let tensor_t = annotation_from_type(&IJType::Tensor(input_number_type))?;
956 match node.operands.len() {
957 2 => Ok(quote! {
958 |x: #tensor_t, y: #tensor_t| x.generalized_contraction(&y, #f_stream_extract, #g_stream).unwrap()
959 }),
960 4 => {
961 let op1 = child_streams[2].clone();
962 let op2 = child_streams[3].clone();
963 Ok(quote! {
964 #op1.generalized_contraction(&#op2, #f_stream_extract, #g_stream).unwrap()
965 })
966 }
967 _ => {
968 panic!(
969 "Expected 2 or 4 operands for GeneralizedContraction, found {}",
970 node.operands.len()
971 );
972 }
973 }
974 }
975}
976
977struct TensorBuilder;
978impl CompileNode for TensorBuilder {
979 fn compile(
980 node: Rc<Node>,
981 _compiler: &mut CompilerContext,
982 child_streams: HashMap<usize, TokenStream>,
983 ) -> Result<TokenStream> {
984 if let Operation::TensorBuilder(builder_name) = &node.op {
985 let builder_name_stream = syn::parse_str::<proc_macro2::TokenStream>(
986 &builder_name.to_string(),
987 )
988 .map_err(|_| {
989 syn::Error::new_spanned(
990 builder_name.to_string(),
991 "Failed to parse builder name into a Rust TokenStream",
992 )
993 })?;
994
995 match node.operands.len() {
996 1 => {
997 let child_stream = child_streams.get(&node.operands[0].id).unwrap().clone();
998 let number_type = node.output_type.extract_number_type().unwrap();
999 let tensor_t = annotation_from_type(&IJType::Tensor(number_type))?;
1000 Ok(quote! {
1001 #tensor_t::#builder_name_stream(#child_stream.to_vec().as_slice())
1002 })
1003 }
1004 0 => {
1005 let tensor_type = node.output_type.extract_signature().unwrap().output;
1006 let tensor_t = annotation_from_type(&tensor_type)?;
1007 Ok(quote! {
1008 |_x: #tensor_t| #tensor_t::#builder_name_stream(_x.to_vec().as_slice())
1009 })
1010 }
1011
1012 _ => {
1013 panic!(
1014 "Expected 0 or 1 operand for TensorBuilder, found {}",
1015 node.operands.len()
1016 );
1017 }
1018 }
1019 } else {
1020 unreachable!()
1021 }
1022 }
1023}
1024
1025struct Transpose;
1026impl CompileNode for Transpose {
1027 fn compile(
1028 node: Rc<Node>,
1029 _compiler: &mut CompilerContext,
1030 child_streams: HashMap<usize, TokenStream>,
1031 ) -> Result<TokenStream> {
1032 match node.operands.len() {
1033 1 => {
1034 let child_stream = child_streams.get(&node.operands[0].id).unwrap().clone();
1035 Ok(quote! {
1036 #child_stream.transpose()
1037 })
1038 }
1039 0 => {
1040 let tensor_type = node.output_type.extract_signature().unwrap().output;
1041 let tensor_t = annotation_from_type(&tensor_type)?;
1042 Ok(quote! {
1043 |_x: #tensor_t| _x.transpose()
1044 })
1045 }
1046 _ => {
1047 panic!(
1048 "Expected 1 operand for Transpose, found {}",
1049 node.operands.len()
1050 );
1051 }
1052 }
1053 }
1054}
1055
1056struct Shape;
1057impl CompileNode for Shape {
1058 fn compile(
1059 node: Rc<Node>,
1060 _compiler: &mut CompilerContext,
1061 child_streams: HashMap<usize, TokenStream>,
1062 ) -> Result<TokenStream> {
1063 match node.operands.len() {
1064 1 => {
1065 let child_stream = child_streams.get(&node.operands[0].id).unwrap().clone();
1066 Ok(quote! {
1067 ijzer_lib::tensor::Tensor::<usize>::from_vec(#child_stream.shape().to_vec(), None)
1068 })
1069 }
1070 0 => {
1071 let tensor_type = node.output_type.extract_signature().unwrap().input[0].clone();
1072 let tensor_t = annotation_from_type(&tensor_type)?;
1073 Ok(quote! {
1074 |_x: #tensor_t| ijzer_lib::tensor::Tensor::<usize>::from_vec(_x.shape().to_vec(), None)
1075 })
1076 }
1077 _ => {
1078 panic!(
1079 "Expected 1 operand for Shape, found {}",
1080 node.operands.len()
1081 );
1082 }
1083 }
1084 }
1085}
1086
1087struct QR;
1088impl CompileNode for QR {
1089 fn compile(
1090 node: Rc<Node>,
1091 _compiler: &mut CompilerContext,
1092 child_streams: HashMap<usize, TokenStream>,
1093 ) -> Result<TokenStream> {
1094 match node.operands.len() {
1095 1 => {
1096 let child_stream = child_streams.get(&node.operands[0].id).unwrap().clone();
1097 Ok(quote! {
1098 #child_stream.qr().unwrap()
1099 })
1100 }
1101 0 => {
1102 let tensor_type = node.output_type.extract_signature().unwrap().input[0].clone();
1103 let tensor_t = annotation_from_type(&tensor_type)?;
1104 Ok(quote! {
1105 |_x: #tensor_t| _x.qr().unwrap()
1106 })
1107 }
1108 _ => {
1109 panic!(
1110 "Expected 0 or 1 operand for QR, found {}",
1111 node.operands.len()
1112 );
1113 }
1114 }
1115 }
1116}
1117
1118struct Svd;
1119impl CompileNode for Svd {
1120 fn compile(
1121 node: Rc<Node>,
1122 _compiler: &mut CompilerContext,
1123 child_streams: HashMap<usize, TokenStream>,
1124 ) -> Result<TokenStream> {
1125 match node.operands.len() {
1126 1 => {
1127 let child_stream = child_streams.get(&node.operands[0].id).unwrap().clone();
1128 Ok(quote! {
1129 #child_stream.svd().unwrap()
1130 })
1131 }
1132 0 => {
1133 let tensor_type = node.output_type.extract_signature().unwrap().input[0].clone();
1134 let tensor_t = annotation_from_type(&tensor_type)?;
1135 Ok(quote! {
1136 |_x: #tensor_t| _x.svd().unwrap()
1137 })
1138 }
1139 _ => {
1140 panic!(
1141 "Expected 0 or 1 operand for SVD, found {}",
1142 node.operands.len()
1143 );
1144 }
1145 }
1146 }
1147}
1148
1149struct Solve;
1150impl CompileNode for Solve {
1151 fn compile(
1152 node: Rc<Node>,
1153 _compiler: &mut CompilerContext,
1154 child_streams: HashMap<usize, TokenStream>,
1155 ) -> Result<TokenStream> {
1156 match node.operands.len() {
1157 2 => {
1158 let lhs_stream = child_streams.get(&node.operands[0].id).unwrap().clone();
1159 let rhs_stream = child_streams.get(&node.operands[1].id).unwrap().clone();
1160 Ok(quote! {
1161 #lhs_stream.solve(&#rhs_stream).unwrap()
1162 })
1163 }
1164 0 => {
1165 let tensor_type0 = node.output_type.extract_signature().unwrap().input[0].clone();
1166 let tensor_type1 = node.output_type.extract_signature().unwrap().input[1].clone();
1167 let tensor_t0 = annotation_from_type(&tensor_type0)?;
1168 let tensor_t1 = annotation_from_type(&tensor_type1)?;
1169 Ok(quote! {
1170 |_x: #tensor_t0, _y: #tensor_t1| _x.solve(&_y).unwrap()
1171 })
1172 }
1173 _ => {
1174 panic!(
1175 "Expected 0 or 2 operands for Solve, found {}",
1176 node.operands.len()
1177 );
1178 }
1179 }
1180 }
1181}
1182
1183struct Diag;
1184impl CompileNode for Diag {
1185 fn compile(
1186 node: Rc<Node>,
1187 _compiler: &mut CompilerContext,
1188 child_streams: HashMap<usize, TokenStream>,
1189 ) -> Result<TokenStream> {
1190 match node.operands.len() {
1191 1 => {
1192 let child_stream = child_streams.get(&node.operands[0].id).unwrap().clone();
1193 let tensor_type = node.output_type.clone();
1194 let tensor_t = annotation_from_type(&tensor_type)?;
1195 Ok(quote! {
1196 #tensor_t::diag(&#child_stream)
1197 })
1198 }
1199 0 => {
1200 let tensor_type = node.output_type.extract_signature().unwrap().input[0].clone();
1201 let tensor_t = annotation_from_type(&tensor_type)?;
1202 Ok(quote! {
1203 |_x: #tensor_t| #tensor_t::diag(&_x)
1204 })
1205 }
1206 _ => {
1207 panic!("Expected 1 operand for Diag, found {}", node.operands.len());
1208 }
1209 }
1210 }
1211}
1212
1213struct Index;
1214impl CompileNode for Index {
1215 fn compile(
1216 node: Rc<Node>,
1217 _compiler: &mut CompilerContext,
1218 child_streams: HashMap<usize, TokenStream>,
1219 ) -> Result<TokenStream> {
1220 let tensor_operand = node.operands[0].clone();
1221 let tensor_stream = child_streams.get(&tensor_operand.id).unwrap().clone();
1222
1223 let index_operands = &node.operands[1..];
1224 let all_numbers = index_operands.iter().all(|n| {
1225 n.output_type
1226 .type_match(&IJType::Number(Some("usize".to_string())))
1227 });
1228 let contains_colon = index_operands
1229 .iter()
1230 .any(|n| n.output_type.type_match(&IJType::Void));
1231
1232 match (all_numbers, contains_colon) {
1233 (true, _) => {
1234 let index_operands_stream = index_operands
1235 .iter()
1236 .map(|n| child_streams.get(&n.id).unwrap().clone())
1237 .collect::<Vec<TokenStream>>();
1238 Ok(quote! {
1239 #tensor_stream[&vec![#(#index_operands_stream),*]].clone()
1240 })
1241 }
1242 (false, true) => {
1243 let mut index_operands_streams = vec![];
1244 for index_operand in index_operands {
1245 match index_operand.output_type {
1246 IJType::Number(_) => {
1247 let stream = child_streams.get(&index_operand.id).unwrap().clone();
1248 index_operands_streams.push(quote! {Some(#stream)});
1249 }
1250 IJType::Void => index_operands_streams.push(quote! {None}),
1251 _ => unreachable!(),
1252 }
1253 }
1254 Ok(quote! {
1255 #tensor_stream.sub_tensor(vec![#(#index_operands_streams),*]).unwrap()
1256 })
1257 }
1258 (false, false) => {
1259 let index_operands_streams = index_operands
1260 .iter()
1261 .map(|n| child_streams.get(&n.id).unwrap().clone())
1262 .collect::<Vec<TokenStream>>();
1263 Ok(quote! {
1264 #tensor_stream.multi_index(vec![#(#index_operands_streams),*]).unwrap()
1265 })
1266 }
1267 }
1268 }
1269}
1270
1271struct Reshape;
1272impl CompileNode for Reshape {
1273 fn compile(
1274 node: Rc<Node>,
1275 compiler: &mut CompilerContext,
1276 child_streams: HashMap<usize, TokenStream>,
1277 ) -> Result<TokenStream> {
1278 match node.output_type.clone() {
1279 IJType::Tensor(_) => {
1280 let tensor_operand = node.operands[0].clone();
1281 let tensor_stream = child_streams.get(&tensor_operand.id).unwrap().clone();
1282
1283 let shape_operand = node.operands[1].clone();
1284 let shape_stream = child_streams.get(&shape_operand.id).unwrap().clone();
1285
1286 let ident = compiler.get_varname(node.id);
1287 Ok(quote! {
1288 { let mut #ident = #tensor_stream; #ident.reshape(&#shape_stream.to_vec()).unwrap(); #ident }
1289 })
1290 }
1291 IJType::Function(signature) => {
1292 let input_types = signature.input.clone();
1293 let input_1_annot = annotation_from_type(&input_types[0])?;
1294 let input_2_annot = annotation_from_type(&input_types[1])?;
1295
1296 let ident = compiler.get_varname(node.id);
1297 Ok(quote! {
1298 |_x: #input_1_annot, _s: #input_2_annot| {
1299 let mut #ident = _x.clone();
1300 #ident.reshape(&_s.to_vec()).unwrap();
1301 #ident
1302 }
1303 })
1304 }
1305 _ => unreachable!(),
1306 }
1307 }
1308}
1309
1310struct UnaryFunction;
1311impl CompileNode for UnaryFunction {
1312 fn compile(
1313 node: Rc<Node>,
1314 _compiler: &mut CompilerContext,
1315 child_streams: HashMap<usize, TokenStream>,
1316 ) -> Result<TokenStream> {
1317 if let Operation::UnaryFunction(function_name) = &node.op {
1318 let function_name_stream = syn::parse_str::<proc_macro2::TokenStream>(
1319 &function_name.to_string(),
1320 )
1321 .map_err(|_| {
1322 syn::Error::new_spanned(
1323 function_name.to_string(),
1324 "Failed to parse function name into a Rust TokenStream",
1325 )
1326 })?;
1327
1328 match node.output_type.clone() {
1329 IJType::Tensor(_) => {
1330 let child_stream = child_streams.get(&node.operands[0].id).unwrap().clone();
1331 Ok(quote! {
1332 #child_stream.map(|_x| _x.#function_name_stream())
1333 })
1334 }
1335 IJType::Number(_) => {
1336 let child_stream = child_streams.get(&node.operands[0].id).unwrap().clone();
1337 Ok(quote! {
1338 #child_stream.#function_name_stream()
1339 })
1340 }
1341 IJType::Function(signature) => match *signature.output.clone() {
1342 IJType::Tensor(_) => {
1343 let input_type = signature.input[0].clone();
1344 let input_annot = annotation_from_type(&input_type)?;
1345 Ok(quote! {
1346 |_x: #input_annot| _x.map(|_y| _y.#function_name_stream())
1347 })
1348 }
1349 IJType::Number(_) => {
1350 let input_type = signature.input[0].clone();
1351 let input_annot = annotation_from_type(&input_type)?;
1352 Ok(quote! {
1353 |_x: #input_annot| _x.#function_name_stream()
1354 })
1355 }
1356 _ => unreachable!(),
1357 },
1358 _ => unreachable!(),
1359 }
1360 } else {
1361 unreachable!()
1362 }
1363 }
1364}
1365
1366struct Range;
1367impl CompileNode for Range {
1368 fn compile(
1369 node: Rc<Node>,
1370 _compiler: &mut CompilerContext,
1371 child_streams: HashMap<usize, TokenStream>,
1372 ) -> Result<TokenStream> {
1373 match node.output_type.clone() {
1374 IJType::Tensor(_) => {
1375 let start_stream = child_streams.get(&node.operands[0].id).unwrap().clone();
1376 let end_stream = child_streams.get(&node.operands[1].id).unwrap().clone();
1377 let tensor_annot = annotation_from_type(&node.output_type)?;
1378 Ok(quote! {
1379 #tensor_annot::range(#start_stream, #end_stream)
1380 })
1381 }
1382 IJType::Function(signature) => {
1383 let input_type = signature.input[0].clone();
1384 let input_annot = annotation_from_type(&input_type)?;
1385 let output_type = signature.output.clone();
1386 let output_annot = annotation_from_type(&output_type)?;
1387 Ok(quote! {
1388 |_start: #input_annot, _end: #input_annot| #output_annot::range(_start, _end)
1389 })
1390 }
1391 _ => unreachable!(),
1392 }
1393 }
1394}
1395
1396#[cfg(test)]
1397mod tests {
1398 use super::*;
1399 use crate::compile;
1400
1401 fn compiler_compare(input: &str, expected: &str) {
1402 let compiled_stream_result = compile(input);
1403 if let Err(e) = &compiled_stream_result {
1404 println!("Error: {:?}", e);
1405 panic!("Compilation failed");
1406 }
1407 let compiled_stream = compiled_stream_result.unwrap();
1408 println!("------------------------");
1409 println!("Input:");
1410 println!("{}", input);
1411 println!("Compiled stream (string):");
1412 println!("'{}'", compiled_stream);
1413
1414 let expected_stream_result = syn::parse_str(expected);
1415 assert!(expected_stream_result.is_ok());
1416 let expected_stream: TokenStream = expected_stream_result.unwrap();
1417
1418 println!("------------------------");
1419 println!("Expected stream (string):");
1420 println!("'{}'", expected_stream);
1421 println!("------------------------");
1422
1423 assert_eq!(
1424 compiled_stream.to_string().replace(" ", ""),
1425 expected_stream.to_string().replace(" ", "")
1426 );
1427 }
1428
1429 #[test]
1430 fn test_empty_input() {
1431 compiler_compare("", "");
1432 }
1433
1434 #[test]
1435 fn test_var() {
1436 compiler_compare("var x: T", "");
1437 }
1438
1439 #[test]
1440 fn test_assign_tensor() {
1441 let input1 = "x: T<_> = [1]<_>";
1442 let input2 = "x = [1]<_>";
1443 let expected = "let x : ijzer_lib:: tensor :: Tensor :: < _ > = ijzer_lib:: tensor :: Tensor :: < _ > :: from_vec (vec ! [1 ] , None) ;";
1444 compiler_compare(input1, expected);
1445 compiler_compare(input2, expected);
1446 }
1447
1448 #[test]
1449 fn test_assign_scalar() {
1450 let input1 = "x: N = 1";
1451 let input2 = "x = 1";
1452 let expected = "let x: _ = 1;";
1453 compiler_compare(input1, expected);
1454 compiler_compare(input2, expected);
1455 }
1456
1457 #[test]
1458 fn test_multiline() {
1459 let input1 = "x = [1]
1460 y = + x x
1461 ";
1462 let input2 = "x = [1]; y= + x x";
1463 let input3 = "
1464 x = [1];
1465 y= + x x;";
1466 let expexted = "let x : ijzer_lib:: tensor :: Tensor :: < _ > = ijzer_lib:: tensor :: Tensor :: < _ > :: from_vec (vec ! [1] , None) ; let y : ijzer_lib:: tensor :: Tensor :: < _ > = x.clone() . apply_binary_op (& x.clone() , | _a : _ , _b : _ | _a + _b) . unwrap () ;";
1467 compiler_compare(input1, expexted);
1468 compiler_compare(input2, expexted);
1469 compiler_compare(input3, expexted);
1470 }
1471
1472 #[test]
1473 fn test_simple_function() {
1474 let input1 = "var y: T; f($x) -> T = + y $x";
1475 let input2 = "var y: T; f($x) = + y $x";
1476 let expexted = "let f = { | __x: ijzer_lib::tensor::Tensor::<_> | y.clone().apply_binary_op(&__x, |_a: _, _b: _| _a + _b).unwrap() } ;";
1477 compiler_compare(input1, expexted);
1478 compiler_compare(input2, expexted);
1479 }
1480
1481 #[test]
1482 fn test_function_apply() {
1483 let input1 = "var x: Fn(N->T); x 1";
1484 let input2 = "var x: Fn(N->T)
1485 x 1";
1486 let expexted = "x (1)";
1487 compiler_compare(input1, expexted);
1488 compiler_compare(input2, expexted);
1489 }
1490
1491 #[test]
1492 fn test_add() {
1493 compiler_compare("+ [1] [2]", "ijzer_lib:: tensor :: Tensor :: < _ > :: from_vec (vec ! [1] , None) . apply_binary_op (& ijzer_lib:: tensor :: Tensor :: < _ > :: from_vec (vec ! [2] , None) , | _a : _ , _b : _ | _a + _b) . unwrap ()");
1494 compiler_compare("+ [1] 2","ijzer_lib:: tensor :: Tensor :: < _ > :: from_vec (vec ! [1] , None) . map (| _x : _ | _x + 2)");
1495 compiler_compare("+ 1 [2]","ijzer_lib:: tensor :: Tensor :: < _ > :: from_vec (vec ! [2] , None) . map (| _x : _ | 1 + _x)");
1496 compiler_compare("+ 1 2", "1+2");
1497 }
1498 #[test]
1499 fn test_subtract() {
1500 compiler_compare("- [1] [2]", "ijzer_lib::tensor::Tensor::<_>::from_vec(vec![1], None).apply_binary_op(&ijzer_lib::tensor::Tensor::<_>::from_vec(vec![2], None), |_a: _, _b: _| _a - _b).unwrap()");
1501 compiler_compare("- [1] 2","ijzer_lib:: tensor :: Tensor :: < _ > :: from_vec (vec ! [1] , None) . map (| _x : _ | _x - 2)");
1502 compiler_compare("- 1 [2]","ijzer_lib:: tensor :: Tensor :: < _ > :: from_vec (vec ! [2] , None) . map (| _x : _ | 1-_x)");
1503 compiler_compare("- 1 2", "1-2");
1504 }
1505 #[test]
1506 fn test_multiply() {
1507 compiler_compare("* [1] [2]", "ijzer_lib::tensor::Tensor::<_>::from_vec(vec![1], None).apply_binary_op(&ijzer_lib::tensor::Tensor::<_>::from_vec(vec![2], None), |_a: _, _b: _| _a * _b).unwrap()");
1508 compiler_compare("* [1] 2","ijzer_lib:: tensor :: Tensor :: < _ > :: from_vec (vec ! [1] , None) . map (| _x : _ | _x * 2)");
1509 compiler_compare("* 1 [2]","ijzer_lib:: tensor :: Tensor :: < _ > :: from_vec (vec ! [2] , None) . map (| _x : _ | 1 * _x)");
1510 compiler_compare("* 1 2", "1*2");
1511 }
1512
1513 #[test]
1514 fn test_negate() {
1515 compiler_compare(
1516 "-[1]",
1517 "ijzer_lib::tensor::Tensor::<_>::from_vec(vec![1], None).map(|a: _| -a)",
1518 );
1519 compiler_compare("var x: T; -x", "x.clone().map(|a: _| -a)");
1520 compiler_compare("var x: N<f64>; -x", "-x");
1521 }
1522
1523 #[test]
1524 fn test_reduction() {
1525 compiler_compare(
1526 "var f: Fn(N,N->N); /f [1,2]",
1527 "ijzer_lib::tensor::Tensor::<_>::from_vec(vec![1,2], None).reduce(f)",
1528 );
1529 compiler_compare(
1530 "/+ [1,2]",
1531 "ijzer_lib::tensor::Tensor::<_>::from_vec(vec![1,2], None).reduce(|a: _, b: _| a + b)",
1532 );
1533
1534 compiler_compare(
1535 "/+ [1,2]<f64>",
1536 "ijzer_lib::tensor::Tensor::<f64>::from_vec(vec![1,2], None).reduce(|a: _, b: _| a + b)",
1537 );
1538 }
1539
1540 #[test]
1541 fn test_lambda_variable() {
1542 let expected1 =
1543 "let f = { | __x: ijzer_lib::tensor::Tensor::<_>| __x . apply_binary_op (& __x , | _a: _ , _b: _ | _a + _b) . unwrap () } ;";
1544 compiler_compare("f($x) = + $x $x", expected1);
1545
1546 let expected2 =
1547 "let f = { | __x: ijzer_lib::tensor::Tensor::<_> , __y: ijzer_lib::tensor::Tensor::<_> | __x . apply_binary_op (& __y , | _a: _ , _b: _ | _a + _b) . unwrap () } ;";
1548 compiler_compare("f($x, $y) = + $x $y", expected2);
1549
1550 let expected3 = "let h = { | __x: ijzer_lib::tensor::Tensor::<_> | k (__x , ijzer_lib::tensor::Tensor::<_>::from_vec(vec![1,2], None)) } ;";
1551 compiler_compare("var k: Fn(T,T->T); h($x) = k $x [1,2]", expected3);
1552 }
1553
1554 #[test]
1555 fn test_function_composition() {
1556 let input = "@(-,+) 1 2";
1557 let expected = "(| _12_1 : _ , _12_2 : _ | (| a : _ | - a) ((| a : _ , b : _ | a + b) (_12_1 , _12_2))) (1 , 2)";
1558 compiler_compare(input, expected);
1559
1560 let input = "@(+) 1 2";
1561 let expected = "(| _6_1 : _ , _6_2 : _ | (| a : _ , b : _ | a + b) (_6_1 , _6_2)) (1 , 2)";
1562 compiler_compare(input, expected);
1563
1564 let input = "var f: Fn(T->T); @(f,-) [1]";
1565 let expected = "(| _10_1 : ijzer_lib:: tensor :: Tensor :: < _ > | (f) ((| x : ijzer_lib:: tensor :: Tensor :: < _ > | x . map (| a : _ | - a)) (_10_1))) (ijzer_lib:: tensor :: Tensor :: < _ > :: from_vec (vec ! [1] , None))";
1566 compiler_compare(input, expected);
1567 }
1568
1569 #[test]
1570 fn test_number_type_from_string() {
1571 let tokens = number_type_from_string("_").unwrap();
1572 assert_eq!(tokens.to_string(), (quote! {_}).to_string());
1573
1574 let tokens = number_type_from_string("f64").unwrap();
1575 assert_eq!(tokens.to_string(), (quote! {f64}).to_string());
1576 }
1577
1578 #[test]
1579 fn test_type_conversion() {
1580 let input = "var x: N; <-T x";
1581 let expected = "ijzer_lib:: tensor :: Tensor :: < _ > :: scalar (x)";
1582 compiler_compare(input, expected);
1583
1584 let input = "var f: Fn(T->T); <-Fn(N->T) f 1";
1585 let expected =
1586 "((| _2_1 : _ | ((f) (ijzer_lib:: tensor :: Tensor :: < _ > :: scalar (_2_1))))) (1)";
1587 compiler_compare(input, expected);
1588
1589 let input = "var f: Fn(T->N); <-Fn(N->T) f 1";
1590 let expected = "((| _2_1 : _ | (ijzer_lib:: tensor :: Tensor :: < _ > :: scalar ((f) (ijzer_lib:: tensor :: Tensor :: < _ > :: scalar (_2_1)))))) (1)";
1591 compiler_compare(input, expected);
1592
1593 let input = "var x: N; <-N x";
1594 let expected = "x";
1595 compiler_compare(input, expected);
1596
1597 let input = "var x: N<f64>; <-N<f64> x";
1598 let expected = "x";
1599 compiler_compare(input, expected);
1600 }
1601
1602 #[test]
1603 fn test_as_function() {
1604 let input = "var f: Fn(T->T); ~f";
1605 let expected = "f";
1606 compiler_compare(input, expected);
1607
1608 let input = "var f: Fn(T<f64>->T<f64>); ~f";
1609 let expected = "f";
1610 compiler_compare(input, expected);
1611
1612 let input = "~+: Fn(N,N->N)";
1613 let expected = "| a : _ , b: _ | a + b";
1614 compiler_compare(input, expected);
1615 }
1616
1617 #[test]
1618 fn test_type_conversion_with_as_function() {
1619 let input = "~(<-Fn(N,N->T) +)";
1620 let expected = "(| _4_1 : _ , _4_2 : _ | (ijzer_lib:: tensor :: Tensor :: < _ > :: scalar ((| a : _ , b : _ | a + b) (_4_1 , _4_2))))";
1621 compiler_compare(input, expected);
1622 }
1623 #[test]
1624 fn test_function_composition_functional() {
1625 let input = "~@(-,+):Fn(N,N->N)";
1626 let expected =
1627 "| _14_1 : _ , _14_2 : _ | (| a : _ | - a) ((| a : _ , b : _ | a + b) (_14_1 , _14_2))";
1628 compiler_compare(input, expected);
1629 }
1630
1631 #[test]
1632 fn test_reduction_in_composition() {
1633 let input = "~@(/+,+):Fn(T,T->N)";
1634 let expected = "| _9_1 : ijzer_lib:: tensor :: Tensor :: < _ > , _9_2 : ijzer_lib:: tensor :: Tensor :: < _ > | (| _1 : ijzer_lib:: tensor :: Tensor :: < _ > | _1 . reduce (| a : _ , b : _ | a + b)) ((| x1 : ijzer_lib:: tensor :: Tensor :: < _ > , x2 : ijzer_lib:: tensor :: Tensor :: < _ > | x1 . apply_binary_op (& x2 , | a : _ , b : _ | a + b) . unwrap ()) (_9_1 , _9_2))";
1635 compiler_compare(input, expected);
1636 }
1637
1638 #[test]
1639 fn test_apply() -> Result<()> {
1640 let input = "var f: Fn(N,N->N); .~f 1 2";
1641 let expected = "(f) (1,2)";
1642 compiler_compare(input, expected);
1643
1644 let input = "var f: Fn(N<f64>,N<f64>->N<f64>); .~f 1<f64> 2<f64>";
1645 let expected = "(f) (1 as f64, 2 as f64)";
1646 compiler_compare(input, expected);
1647
1648 let input = "var f: Fn(N->Fn(N->N)); .(f 1) 2";
1649 let expected = "(f (1)) (2)";
1650 compiler_compare(input, expected);
1651 Ok(())
1652 }
1653
1654 #[test]
1655 fn test_lambda_variable_functional() -> Result<()> {
1656 let input = "g($x:Fn(N,N->N)) -> N = /$x [1]";
1657 let expected = "let g = { | __x : fn (_ , _) -> _ | ijzer_lib:: tensor :: Tensor :: < _ > :: from_vec (vec ! [1] , None) . reduce (__x) } ;";
1658 compiler_compare(input, expected);
1659
1660 let input = "g($x:Fn(N->T), $y:Fn(T->N)) -> Fn(T->T) = ~@($x,$y)";
1661 let expected = "let g = { | __x : fn (_) -> ijzer_lib:: tensor :: Tensor :: < _ > , __y : fn (ijzer_lib:: tensor :: Tensor :: < _ >) -> _ | | _4_1 : ijzer_lib:: tensor :: Tensor :: < _ > | (__x) ((__y) (_4_1)) } ;";
1662 compiler_compare(input, expected);
1663
1664 Ok(())
1665 }
1666
1667 #[test]
1668 fn test_generalized_contraction() -> Result<()> {
1669 let input = "?/+* [1] [2]";
1670 let expected = "ijzer_lib:: tensor :: Tensor :: < _ > :: from_vec (vec ! [1] , None) . generalized_contraction (& ijzer_lib:: tensor :: Tensor :: < _ > :: from_vec (vec ! [2] , None) , (| z : & ijzer_lib:: tensor :: Tensor :: < _ > | (| _1 : ijzer_lib:: tensor :: Tensor :: < _ > | _1 . reduce (| a : _ , b : _ | a + b)) (z . clone ())) , | a : _ , b : _ | a * b) . unwrap ()";
1671 compiler_compare(input, expected);
1672
1673 Ok(())
1674 }
1675
1676 #[test]
1677 fn test_generalized_contraction_functional() -> Result<()> {
1678 let input = "~?/+*";
1679 let expected = "| x : ijzer_lib:: tensor :: Tensor :: < _ > , y : ijzer_lib:: tensor :: Tensor :: < _ > | x . generalized_contraction (& y , (| z : & ijzer_lib:: tensor :: Tensor :: < _ > | (| _1 : ijzer_lib:: tensor :: Tensor :: < _ > | _1 . reduce (| a : _ , b : _ | a + b)) (z . clone ())) , | a : _ , b : _ | a * b) . unwrap ()";
1680 compiler_compare(input, expected);
1681
1682 Ok(())
1683 }
1684
1685 #[test]
1686 fn test_tensor_builder() -> Result<()> {
1687 let input = "zeros [1,2]";
1688 let expected = "ijzer_lib::tensor::Tensor::<_>::zeros(ijzer_lib::tensor::Tensor::<_>::from_vec(vec![1,2],None).to_vec().as_slice())";
1689 compiler_compare(input, expected);
1690
1691 let input = "ones<f64> [1]";
1692 let expected = "ijzer_lib::tensor::Tensor::<f64>::ones(ijzer_lib::tensor::Tensor::<_>::from_vec(vec![1],None).to_vec().as_slice())";
1693 compiler_compare(input, expected);
1694
1695 let input = "var x: T<usize>; randu<f32> x";
1696 let expected = "ijzer_lib::tensor::Tensor::<f32>::randu(x.clone().to_vec().as_slice())";
1697 compiler_compare(input, expected);
1698
1699 let input = "var x: T<usize>; randn<f32> x";
1700 let expected = "ijzer_lib::tensor::Tensor::<f32>::randn(x.clone().to_vec().as_slice())";
1701 compiler_compare(input, expected);
1702
1703 let input = "eye<f32> [3,3]<usize>";
1704 let expected = "ijzer_lib::tensor::Tensor::<f32>::eye(ijzer_lib::tensor::Tensor::<usize>::from_vec(vec![3,3],None).to_vec().as_slice())";
1705 compiler_compare(input, expected);
1706
1707 let input = "var x: T<usize>; .~eye x";
1708 let expected = "(|_x: ijzer_lib::tensor::Tensor::<_>| ijzer_lib::tensor::Tensor::<_>::eye(_x.to_vec().as_slice()))(x.clone())";
1709 compiler_compare(input, expected);
1710
1711 Ok(())
1712 }
1713
1714 #[test]
1715 fn test_transpose() -> Result<()> {
1716 let input = "var x: T<f64>; |x";
1717 let expected = "x.clone().transpose()";
1718 compiler_compare(input, expected);
1719
1720 let input = "var x: T<f64>; .~|x";
1721 let expected = "(|_x: ijzer_lib::tensor::Tensor::<_>| _x.transpose())(x.clone())";
1722 compiler_compare(input, expected);
1723
1724 Ok(())
1725 }
1726
1727 #[test]
1728 fn test_shape() -> Result<()> {
1729 let input = "var x: T; %x";
1730 let expected =
1731 "ijzer_lib::tensor::Tensor::<usize>::from_vec(x.clone().shape().to_vec(), None)";
1732 compiler_compare(input, expected);
1733
1734 let input = "var x: T;.~%x";
1735 let expected = "(|_x: ijzer_lib::tensor::Tensor::<_>| ijzer_lib::tensor::Tensor::<usize>::from_vec(_x.shape().to_vec(), None))(x.clone())";
1736 compiler_compare(input, expected);
1737
1738 Ok(())
1739 }
1740
1741 #[test]
1742 fn test_group_assign() -> Result<()> {
1743 let input = "(x: N<a>, y: N) = (1<a>,2)";
1744 let expected = "let (x , y) : (a, _) = (1 as a , 2) ;";
1745 compiler_compare(input, expected);
1746
1747 let input = "var f: Fn(T->(T,T)); (x: T, y: T) = f [1]";
1748 let expected = "let (x , y) : (ijzer_lib:: tensor :: Tensor :: < _ > , ijzer_lib:: tensor :: Tensor :: < _ >) = f (ijzer_lib:: tensor :: Tensor :: < _ > :: from_vec (vec ! [1] , None)) ;";
1749 compiler_compare(input, expected);
1750 Ok(())
1751 }
1752
1753 #[test]
1754 fn test_solve() -> Result<()> {
1755 let input = r"var x: T; var y: T; \ x y";
1756 let expected = "x.clone().solve(&y.clone()).unwrap()";
1757 compiler_compare(input, expected);
1758
1759 let input = r"~\";
1760 let expected = "| _x: ijzer_lib:: tensor :: Tensor :: < _ > , _y: ijzer_lib:: tensor :: Tensor :: < _ > | _x.solve(&_y).unwrap()";
1761 compiler_compare(input, expected);
1762
1763 let input = r"var x: T; var y: T; z = \ x y";
1764 let expected =
1765 "let z: ijzer_lib:: tensor :: Tensor :: < _ > = x.clone().solve(&y.clone()).unwrap();";
1766 compiler_compare(input, expected);
1767 Ok(())
1768 }
1769
1770 #[test]
1771 fn test_qr() -> Result<()> {
1772 let input = r"var x: T; qr x";
1773 let expected = "x.clone().qr().unwrap()";
1774 compiler_compare(input, expected);
1775
1776 let input = r"~qr ";
1777 let expected = "| _x: ijzer_lib:: tensor :: Tensor :: < _ > | _x.qr().unwrap()";
1778 compiler_compare(input, expected);
1779
1780 let input = r"var x: T; (q,r) = qr x";
1781 let expected = "let (q , r) : (ijzer_lib:: tensor :: Tensor :: < _ > , ijzer_lib:: tensor :: Tensor :: < _ >) = x.clone().qr () . unwrap () ;";
1782 compiler_compare(input, expected);
1783 Ok(())
1784 }
1785
1786 #[test]
1787 fn test_svd() -> Result<()> {
1788 let input = r"var x: T; svd x";
1789 let expected = "x.clone().svd().unwrap()";
1790 compiler_compare(input, expected);
1791
1792 let input = r"~svd";
1793 let expected = "| _x: ijzer_lib:: tensor :: Tensor :: < _ > | _x.svd().unwrap()";
1794 compiler_compare(input, expected);
1795
1796 let input = r"var x: T; (u,s,v) = svd x";
1797 let expected = "let (u , s , v) : (ijzer_lib:: tensor :: Tensor :: < _ > , ijzer_lib:: tensor :: Tensor :: < _ > , ijzer_lib:: tensor :: Tensor :: < _ >) = x.clone().svd () . unwrap () ;";
1798 compiler_compare(input, expected);
1799 Ok(())
1800 }
1801
1802 #[test]
1803 fn test_diag() -> Result<()> {
1804 let input = r"var x: T; diag x";
1805 let expected = "ijzer_lib::tensor::Tensor::<_>::diag(&x.clone())";
1806 compiler_compare(input, expected);
1807
1808 let input = r"~diag";
1809 let expected =
1810 "| _x: ijzer_lib:: tensor :: Tensor :: < _ > | ijzer_lib::tensor::Tensor::<_>::diag(&_x)";
1811 compiler_compare(input, expected);
1812
1813 let input = r"var x: T; y = diag x";
1814 let expected =
1815 "let y: ijzer_lib:: tensor :: Tensor :: < _ > = ijzer_lib::tensor::Tensor::<_>::diag(&x.clone());";
1816 compiler_compare(input, expected);
1817
1818 Ok(())
1819 }
1820
1821 #[test]
1822 fn test_multi_index() -> Result<()> {
1823 let input = r"var x: T; <| x [1,2]";
1824 let expected = "x.clone() [& vec ! [1 , 2]] . clone ()";
1825 compiler_compare(input, expected);
1826
1827 let input = r"var x: T; <| x [1,:]";
1828 let expected = "x.clone().sub_tensor (vec ! [Some(1), None]).unwrap()";
1829 compiler_compare(input, expected);
1830
1831 let input = r"var x: T; var s1: T; var s2: T; <| x [s1,s2]";
1832 let expected = "x.clone().multi_index (vec ! [s1.clone() , s2.clone()]).unwrap()";
1833 compiler_compare(input, expected);
1834
1835 Ok(())
1836 }
1837
1838 #[test]
1839 fn test_identity() -> Result<()> {
1840 let input = r"var x: T; I x";
1841 let expected = "x.clone()";
1842 compiler_compare(input, expected);
1843
1844 let input = r"~I:Fn(T->T)";
1845 let expected = "| _2 | _2";
1846 compiler_compare(input, expected);
1847
1848 Ok(())
1849 }
1850
1851 #[test]
1852 fn test_array_from_arrays() -> Result<()> {
1853 let input = r"var x: T; var y: T; [x,y]";
1854 let expected =
1855 "ijzer_lib::tensor::Tensor::<_>::from_tensors(&[x.clone(), y.clone()]).unwrap()";
1856 compiler_compare(input, expected);
1857
1858 Ok(())
1859 }
1860
1861 #[test]
1862 fn test_binop_functional() -> Result<()> {
1863 let input = r"~+:Fn(T,N->T)";
1864 let expected = "| x : ijzer_lib:: tensor :: Tensor :: < _ > , y : _ | x.map(|a: _| a + y)";
1865 compiler_compare(input, expected);
1866
1867 let input = r"~-:Fn(N,T->T)";
1868 let expected = "| y : _ , x : ijzer_lib:: tensor :: Tensor :: < _ > | x.map(|a: _| y-a)";
1869 compiler_compare(input, expected);
1870
1871 Ok(())
1872 }
1873
1874 #[test]
1875 fn test_single_group_unpacking() -> Result<()> {
1876 let input = r"var x: T; (x)";
1877 let expected = "x.clone()";
1878 compiler_compare(input, expected);
1879
1880 Ok(())
1881 }
1882
1883 #[test]
1884 fn test_reshape() -> Result<()> {
1885 let input = r"var x: T; var s: T<usize>; >% x s";
1886 let expected = "{ let mut _4 = x . clone () ; _4 . reshape (& s . clone () . to_vec ()) . unwrap () ; _4 }";
1887 compiler_compare(input, expected);
1888
1889 let input = r"var x: T; var s: T<usize>; .~>% x s";
1890 let expected = "(| _x : ijzer_lib:: tensor :: Tensor :: < _ > , _s : ijzer_lib:: tensor :: Tensor :: < usize > | { let mut _2 = _x.clone() ; _2 . reshape (& _s . to_vec ()) . unwrap () ; _2 }) (x . clone () , s . clone ())";
1891 compiler_compare(input, expected);
1892
1893 Ok(())
1894 }
1895
1896 #[test]
1897 fn test_unary_functional() -> Result<()> {
1898 let input = r"var x: T; sin x";
1899 let expected = "x.clone().map(|_x| _x.sin())";
1900 compiler_compare(input, expected);
1901
1902 let input = r"var x: N; sin x";
1903 let expected = "x.sin()";
1904 compiler_compare(input, expected);
1905
1906 let input = r"var x: T; .~sin: Fn(T->T) x";
1907 let expected = "(| _x : ijzer_lib :: tensor :: Tensor :: < _ > | _x . map (| _y | _y . sin ())) (x . clone ())";
1908 compiler_compare(input, expected);
1909
1910 let input = r"var x: N; .~sin: Fn(N->N) x";
1911 let expected = "(| _x : _ | _x . sin ()) (x)";
1912 compiler_compare(input, expected);
1913
1914 Ok(())
1915 }
1916
1917 #[test]
1918 fn test_binary_functional() -> Result<()> {
1919 let input = r"var x: T; var y: T; + x y";
1920 let expected = "x.clone().apply_binary_op(&y.clone(), |_a: _, _b: _| _a + _b).unwrap()";
1921 compiler_compare(input, expected);
1922
1923 let input = r"var x: T; var y: T; * x y";
1924 let expected = "x.clone().apply_binary_op(&y.clone(), |_a: _, _b: _| _a * _b).unwrap()";
1925 compiler_compare(input, expected);
1926
1927 let input = r"var x: T; var y: T; /: x y";
1928 let expected = "x.clone().apply_binary_op(&y.clone(), |_a: _, _b: _| _a / _b).unwrap()";
1929 compiler_compare(input, expected);
1930
1931 let input = r"var x: T; var y: T; ^ x y";
1932 let expected = "x.clone().apply_binary_op(&y.clone(), |_a: _, _b: _| _a.pow(_b)).unwrap()";
1933 compiler_compare(input, expected);
1934
1935 let input = r"var x: T; var y: T; max x y";
1936 let expected = "x.clone().apply_binary_op(&y.clone(), |_a: _, _b: _| _a.max(_b)).unwrap()";
1937 compiler_compare(input, expected);
1938
1939 let input = r"var x: T; var y: T; min x y";
1940 let expected = "x.clone().apply_binary_op(&y.clone(), |_a: _, _b: _| _a.min(_b)).unwrap()";
1941 compiler_compare(input, expected);
1942
1943 let input = r"var x: T; var y: T; == x y";
1944 let expected = "x.clone().apply_binary_op(&y.clone(), |_a: _, _b: _| ijzer_lib::comparison_funcs::equals(_a, _b)).unwrap()";
1945 compiler_compare(input, expected);
1946
1947 let input = r"var x: T; var y: T; != x y";
1948 let expected = "x.clone().apply_binary_op(&y.clone(), |_a: _, _b: _| ijzer_lib::comparison_funcs::not_equals(_a, _b)).unwrap()";
1949 compiler_compare(input, expected);
1950
1951 let input = r"var x: T; var y: T; >. x y";
1952 let expected = "x.clone().apply_binary_op(&y.clone(), |_a: _, _b: _| ijzer_lib::comparison_funcs::greater_than(_a, _b)).unwrap()";
1953 compiler_compare(input, expected);
1954
1955 let input = r"var x: T; var y: T; <. x y";
1956 let expected = "x.clone().apply_binary_op(&y.clone(), |_a: _, _b: _| ijzer_lib::comparison_funcs::less_than(_a, _b)).unwrap()";
1957 compiler_compare(input, expected);
1958
1959 let input = r"var x: T; var y: T; >= x y";
1960 let expected = "x.clone().apply_binary_op(&y.clone(), |_a: _, _b: _| ijzer_lib::comparison_funcs::greater_than_or_equal(_a, _b)).unwrap()";
1961 compiler_compare(input, expected);
1962
1963 let input = r"var x: T; var y: T; <= x y";
1964 let expected = "x.clone().apply_binary_op(&y.clone(), |_a: _, _b: _| ijzer_lib::comparison_funcs::less_than_or_equal(_a, _b)).unwrap()";
1965 compiler_compare(input, expected);
1966
1967 let input = r"var x: T; var y: T; && x y";
1968 let expected = "x.clone().apply_binary_op(&y.clone(), |_a: _, _b: _| ijzer_lib::comparison_funcs::and(_a, _b)).unwrap()";
1969 compiler_compare(input, expected);
1970
1971 let input = r"var x: T; var y: T; || x y";
1972 let expected = "x.clone().apply_binary_op(&y.clone(), |_a: _, _b: _| ijzer_lib::comparison_funcs::or(_a, _b)).unwrap()";
1973 compiler_compare(input, expected);
1974
1975 Ok(())
1976 }
1977
1978 #[test]
1979 fn test_range() -> Result<()> {
1980 let input = r".. 1 2 ";
1981 let expected = "ijzer_lib::tensor::Tensor::<_>::range(1, 2)";
1982 compiler_compare(input, expected);
1983
1984 let input = r"var a: N; var b: N; .. a b";
1985 let expected = "ijzer_lib::tensor::Tensor::<_>::range(a, b)";
1986 compiler_compare(input, expected);
1987
1988 let input = r"var a: N<usize>; var b: N; .. a b";
1989 let expected = "ijzer_lib::tensor::Tensor::<usize>::range(a, b)";
1990 compiler_compare(input, expected);
1991
1992 let input = r"var a: N; var b: N<usize>; .. a b";
1993 let expected = "ijzer_lib::tensor::Tensor::<usize>::range(a, b)";
1994 compiler_compare(input, expected);
1995
1996 let input = r"var a: N<usize>; var b: N<usize>; .. a b";
1997 let expected = "ijzer_lib::tensor::Tensor::<usize>::range(a, b)";
1998 compiler_compare(input, expected);
1999
2000 let input = r"~ .. ";
2001 let expected =
2002 "| _start: _ , _end: _ | ijzer_lib::tensor::Tensor::<_>::range(_start, _end)";
2003 compiler_compare(input, expected);
2004 Ok(())
2005 }
2006}