()
| 2057 | |
| 2058 | #[test] |
| 2059 | fn test_resharing() -> Result<()> { |
| 2060 | { |
| 2061 | let c = create_context()?; |
| 2062 | let g = c.create_graph()?; |
| 2063 | let i1 = g.input(array_type(vec![2, 10], BIT))?; |
| 2064 | let i2 = g.input(array_type(vec![10, 3], BIT))?; |
| 2065 | let prod = i1.matmul(i2)?; |
| 2066 | let out = prod.sum(vec![0])?; |
| 2067 | out.set_as_output()?; |
| 2068 | g.finalize()?; |
| 2069 | g.set_as_main()?; |
| 2070 | c.finalize()?; |
| 2071 | |
| 2072 | let shared_nodes = propagate_private_annotations(g.clone(), vec![true, true])?.0; |
| 2073 | let reshared_nodes = get_nodes_to_reshare(&g, &shared_nodes)?; |
| 2074 | |
| 2075 | assert!(reshared_nodes.len() == 1); |
| 2076 | assert!(reshared_nodes.contains(&out)); |
| 2077 | |
| 2078 | let shared_nodes = propagate_private_annotations(g.clone(), vec![false, true])?.0; |
| 2079 | let reshared_nodes = get_nodes_to_reshare(&g, &shared_nodes)?; |
| 2080 | |
| 2081 | assert!(reshared_nodes.len() == 0); |
| 2082 | |
| 2083 | let shared_nodes = propagate_private_annotations(g.clone(), vec![false, false])?.0; |
| 2084 | let reshared_nodes = get_nodes_to_reshare(&g, &shared_nodes)?; |
| 2085 | |
| 2086 | assert!(reshared_nodes.len() == 0); |
| 2087 | } |
| 2088 | |
| 2089 | { |
| 2090 | let c = create_context()?; |
| 2091 | let g = c.create_graph()?; |
| 2092 | let i1 = g.input(array_type(vec![2, 10], BIT))?; |
| 2093 | let i2 = g.input(array_type(vec![10, 3], BIT))?; |
| 2094 | let prod = i1.matmul(i2)?; |
| 2095 | let i3 = g.input(array_type(vec![3, 4], BIT))?; |
| 2096 | let out = prod.matmul(i3)?; |
| 2097 | out.set_as_output()?; |
| 2098 | g.finalize()?; |
| 2099 | g.set_as_main()?; |
| 2100 | c.finalize()?; |
| 2101 | |
| 2102 | let shared_nodes = propagate_private_annotations(g.clone(), vec![true, true, true])?.0; |
| 2103 | let reshared_nodes = get_nodes_to_reshare(&g, &shared_nodes)?; |
| 2104 | |
| 2105 | assert!(reshared_nodes.len() == 2); |
| 2106 | assert!(reshared_nodes.contains(&prod)); |
| 2107 | assert!(reshared_nodes.contains(&out)); |
| 2108 | |
| 2109 | let shared_nodes = propagate_private_annotations(g.clone(), vec![false, true, true])?.0; |
| 2110 | let reshared_nodes = get_nodes_to_reshare(&g, &shared_nodes)?; |
| 2111 | |
| 2112 | assert!(reshared_nodes.len() == 1); |
| 2113 | assert!(reshared_nodes.contains(&out)); |
| 2114 | |
| 2115 | let shared_nodes = propagate_private_annotations(g.clone(), vec![true, true, false])?.0; |
| 2116 | let reshared_nodes = get_nodes_to_reshare(&g, &shared_nodes)?; |
nothing calls this directly
no test coverage detected