()
| 277 | |
| 278 | #[test] |
| 279 | fn test_resharing() -> Result<()> { |
| 280 | { |
| 281 | let c = create_context()?; |
| 282 | let g = c.create_graph()?; |
| 283 | let i1 = g.input(array_type(vec![2, 10], BIT))?; |
| 284 | let i2 = g.input(array_type(vec![10, 3], BIT))?; |
| 285 | let prod = i1.matmul(i2)?; |
| 286 | let out = prod.sum(vec![0])?; |
| 287 | out.set_as_output()?; |
| 288 | g.finalize()?; |
| 289 | g.set_as_main()?; |
| 290 | c.finalize()?; |
| 291 | |
| 292 | let shared_nodes = propagate_private_annotations(g.clone(), vec![true, true])?.0; |
| 293 | let reshared_nodes = get_nodes_to_reshare(&g, &shared_nodes)?; |
| 294 | |
| 295 | assert!(reshared_nodes.len() == 1); |
| 296 | assert!(reshared_nodes.contains(&out)); |
| 297 | |
| 298 | let shared_nodes = propagate_private_annotations(g.clone(), vec![false, true])?.0; |
| 299 | let reshared_nodes = get_nodes_to_reshare(&g, &shared_nodes)?; |
| 300 | |
| 301 | assert!(reshared_nodes.len() == 0); |
| 302 | |
| 303 | let shared_nodes = propagate_private_annotations(g.clone(), vec![false, false])?.0; |
| 304 | let reshared_nodes = get_nodes_to_reshare(&g, &shared_nodes)?; |
| 305 | |
| 306 | assert!(reshared_nodes.len() == 0); |
| 307 | } |
| 308 | |
| 309 | { |
| 310 | let c = create_context()?; |
| 311 | let g = c.create_graph()?; |
| 312 | let i1 = g.input(array_type(vec![2, 10], BIT))?; |
| 313 | let i2 = g.input(array_type(vec![10, 3], BIT))?; |
| 314 | let prod = i1.matmul(i2)?; |
| 315 | let i3 = g.input(array_type(vec![3, 4], BIT))?; |
| 316 | let out = prod.matmul(i3)?; |
| 317 | out.set_as_output()?; |
| 318 | g.finalize()?; |
| 319 | g.set_as_main()?; |
| 320 | c.finalize()?; |
| 321 | |
| 322 | let shared_nodes = propagate_private_annotations(g.clone(), vec![true, true, true])?.0; |
| 323 | let reshared_nodes = get_nodes_to_reshare(&g, &shared_nodes)?; |
| 324 | |
| 325 | assert!(reshared_nodes.len() == 2); |
| 326 | assert!(reshared_nodes.contains(&prod)); |
| 327 | assert!(reshared_nodes.contains(&out)); |
| 328 | |
| 329 | let shared_nodes = propagate_private_annotations(g.clone(), vec![false, true, true])?.0; |
| 330 | let reshared_nodes = get_nodes_to_reshare(&g, &shared_nodes)?; |
| 331 | |
| 332 | assert!(reshared_nodes.len() == 1); |
| 333 | assert!(reshared_nodes.contains(&out)); |
| 334 | |
| 335 | let shared_nodes = propagate_private_annotations(g.clone(), vec![true, true, false])?.0; |
| 336 | let reshared_nodes = get_nodes_to_reshare(&g, &shared_nodes)?; |
nothing calls this directly
no test coverage detected