Skip to content

Commit b5f1674

Browse files
committed
🔧 chore: apply prek fixes
1 parent e047f49 commit b5f1674

2 files changed

Lines changed: 16 additions & 7 deletions

File tree

paddle/fluid/pybind/eager_py_layer.cc

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1049,8 +1049,7 @@ PyObject* pylayer_hold_tensors(PyObject* self_, PyObject* args) {
10491049
// Re-install impl() on any Python Tensor previously registered via
10501050
// _hold_tensors whose impl_ has been nulled by _clear_dataptr(). Typically
10511051
// called at the start of backward before recompute re-runs forward.
1052-
PyObject* pylayer_restore_held_tensors(PyObject* self_,
1053-
PyObject* /*unused*/) {
1052+
PyObject* pylayer_restore_held_tensors(PyObject* self_, PyObject* /*unused*/) {
10541053
EAGER_TRY
10551054
auto* self = reinterpret_cast<PyLayerObject*>(self_);
10561055
if (self->closure_obj && !self->closure_tensor_hold_helper.empty()) {

test/legacy_test/test_pylayer_clear_dataptr.py

Lines changed: 15 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -535,8 +535,12 @@ def run_fn(a, b):
535535
loss.backward()
536536
run_fn(a_ref, b_ref).backward()
537537

538-
np.testing.assert_allclose(a.grad.numpy(), a_ref.grad.numpy(), rtol=1e-5)
539-
np.testing.assert_allclose(b.grad.numpy(), b_ref.grad.numpy(), rtol=1e-5)
538+
np.testing.assert_allclose(
539+
a.grad.numpy(), a_ref.grad.numpy(), rtol=1e-5
540+
)
541+
np.testing.assert_allclose(
542+
b.grad.numpy(), b_ref.grad.numpy(), rtol=1e-5
543+
)
540544

541545
def test_recompute_closure_tensors(self):
542546
"""Closure captures Tensor / tuple / list / dict: all restored."""
@@ -614,8 +618,12 @@ def fn(inp):
614618
ref_fn(inp_ref).backward()
615619

616620
self.assertIsNone(inp.grad)
617-
np.testing.assert_allclose(w1.grad.numpy(), w1_ref.grad.numpy(), rtol=1e-5)
618-
np.testing.assert_allclose(w2.grad.numpy(), w2_ref.grad.numpy(), rtol=1e-5)
621+
np.testing.assert_allclose(
622+
w1.grad.numpy(), w1_ref.grad.numpy(), rtol=1e-5
623+
)
624+
np.testing.assert_allclose(
625+
w2.grad.numpy(), w2_ref.grad.numpy(), rtol=1e-5
626+
)
619627

620628
def test_recompute_layer_forward_closure(self):
621629
"""paddle.nn.Layer branch of _closure_cell_values."""
@@ -648,7 +656,9 @@ def forward(self, x): # pragma: no cover
648656
loss.backward()
649657
layer_ref(x_ref).backward()
650658

651-
np.testing.assert_allclose(x.grad.numpy(), x_ref.grad.numpy(), rtol=1e-5)
659+
np.testing.assert_allclose(
660+
x.grad.numpy(), x_ref.grad.numpy(), rtol=1e-5
661+
)
652662
np.testing.assert_allclose(
653663
bias.grad.numpy(), bias_ref.grad.numpy(), rtol=1e-5
654664
)

0 commit comments

Comments
 (0)