Returns the inputs of op, crossing closure boundaries where necessary. Args: op: Operation xs_set: ObjectIdentitySet of Tensors we are differentiating w.r.t. Returns: A list of tensors. The tensors may be from multiple Graph/FuncGraphs if op is in a FuncGraph and has captured i
(op, xs_set)
| 469 | # TODO(skyewm): plumbing xs through everywhere is ugly, consider making |
| 470 | # _GradientsHelper a class with xs as a member variable. |
| 471 | def _Inputs(op, xs_set): |
| 472 | """Returns the inputs of op, crossing closure boundaries where necessary. |
| 473 | |
| 474 | Args: |
| 475 | op: Operation |
| 476 | xs_set: ObjectIdentitySet of Tensors we are differentiating w.r.t. |
| 477 | |
| 478 | Returns: |
| 479 | A list of tensors. The tensors may be from multiple Graph/FuncGraphs if op |
| 480 | is in a FuncGraph and has captured inputs. |
| 481 | """ |
| 482 | if _IsFunction(op.graph): # pylint: disable=protected-access |
| 483 | inputs = [] |
| 484 | for t in op.inputs: |
| 485 | # If we're differentiating w.r.t. `t`, do not attempt to traverse through |
| 486 | # it to a captured value. The algorithm needs to "see" `t` in this case, |
| 487 | # even if it's a function input for a captured value, whereas usually we'd |
| 488 | # like to traverse through these closures as if the captured value was the |
| 489 | # direct input to op. |
| 490 | if t not in xs_set: |
| 491 | t = _MaybeCaptured(t) |
| 492 | inputs.append(t) |
| 493 | return inputs |
| 494 | else: |
| 495 | return op.inputs |
| 496 | |
| 497 | |
| 498 | def _Consumers(t, func_graphs): |
no test coverage detected