(
self,
cond_fn: torch.fx.GraphModule,
body_fn: torch.fx.GraphModule,
carried_inputs: List[Argument],
additional_inputs: List[Argument],
meta: NodeMetadata,
)
| 539 | ) |
| 540 | |
| 541 | def call_while( |
| 542 | self, |
| 543 | cond_fn: torch.fx.GraphModule, |
| 544 | body_fn: torch.fx.GraphModule, |
| 545 | carried_inputs: List[Argument], |
| 546 | additional_inputs: List[Argument], |
| 547 | meta: NodeMetadata, |
| 548 | ) -> ProxyValue: |
| 549 | cond_fn = self.call_submodule(cond_fn, (*carried_inputs, *additional_inputs)) |
| 550 | body_fn = self.call_submodule(body_fn, (*carried_inputs, *additional_inputs)) |
| 551 | assert cond_fn is not None |
| 552 | assert body_fn is not None |
| 553 | return self._fx( |
| 554 | "call_function", |
| 555 | torch.ops.higher_order.while_loop, |
| 556 | ( |
| 557 | cond_fn.graph_module, |
| 558 | body_fn.graph_module, |
| 559 | carried_inputs, |
| 560 | additional_inputs, |
| 561 | ), |
| 562 | {}, |
| 563 | meta, |
| 564 | ) |
| 565 | |
| 566 | def call_map( |
| 567 | self, |
no test coverage detected