@@ -692,31 +692,52 @@ PyObject* pylayer_method_apply(PyObject* cls,
692692// DenseTensors found (Tensor / Tuple / List, recursively). Used by
693693// tensor_properties_set_container to hold strong references so that
694694// _clear_dataptr() cannot free the underlying allocation before backward.
695- static void CollectDenseTensors (
696- PyObject* obj, std::vector<std::shared_ptr<phi::DenseTensor>>* holder) {
695+ // DFS-walks obj (tuple/list tree) and calls fn(tensor) for every Tensor leaf.
696+ // Both CollectDenseTensors and RestoreDenseTensors are built on top of this.
697+ template <typename Fn>
698+ static void WalkDenseTensors (PyObject* obj, Fn&& fn) {
697699 if (!obj || obj == Py_None) return ;
698700 if (PyCheckTensor (obj)) {
699- const auto & tensor = reinterpret_cast <TensorObject*>(obj)->tensor ;
700- if (tensor.impl () && tensor.is_dense_tensor ()) {
701- holder->push_back (
702- std::static_pointer_cast<phi::DenseTensor>(tensor.impl ()));
703- }
701+ fn (reinterpret_cast <TensorObject*>(obj)->tensor );
704702 return ;
705703 }
706704 if (PyTuple_Check (obj)) {
707705 Py_ssize_t n = PyTuple_GET_SIZE (obj);
708706 for (Py_ssize_t i = 0 ; i < n; ++i)
709- CollectDenseTensors (PyTuple_GET_ITEM (obj, i), holder );
707+ WalkDenseTensors (PyTuple_GET_ITEM (obj, i), fn );
710708 return ;
711709 }
712710 if (PyList_Check (obj)) {
713711 Py_ssize_t n = PyList_GET_SIZE (obj);
714712 for (Py_ssize_t i = 0 ; i < n; ++i)
715- CollectDenseTensors (PyList_GET_ITEM (obj, i), holder );
713+ WalkDenseTensors (PyList_GET_ITEM (obj, i), fn );
716714 return ;
717715 }
718716}
719717
718+ static void CollectDenseTensors (
719+ PyObject* obj, std::vector<std::shared_ptr<phi::DenseTensor>>* holder) {
720+ WalkDenseTensors (obj, [holder](const paddle::Tensor& tensor) {
721+ if (tensor.impl () && tensor.is_dense_tensor ())
722+ holder->push_back (
723+ std::static_pointer_cast<phi::DenseTensor>(tensor.impl ()));
724+ });
725+ }
726+
727+ // Re-installs impl() for tensors cleared by _clear_dataptr(), using the
728+ // shared_ptrs stored in holder (same DFS order as CollectDenseTensors).
729+ static void RestoreDenseTensors (
730+ PyObject* obj,
731+ const std::vector<std::shared_ptr<phi::DenseTensor>>& holder) {
732+ size_t idx = 0 ;
733+ WalkDenseTensors (obj, [&holder, &idx](paddle::Tensor& tensor) {
734+ if (idx < holder.size ()) {
735+ if (!tensor.impl ()) tensor.set_impl (holder[idx]);
736+ ++idx;
737+ }
738+ });
739+ }
740+
720741PyObject* call_unpack_hook (PyLayerObject* self) {
721742 auto unpack_hook = self->unpack_hook ;
722743 auto packed_value = self->container ;
@@ -767,45 +788,13 @@ PyObject* tensor_properties_get_container(PyLayerObject* self, void* closure) {
767788 if (self->container_be_packed ) {
768789 return call_unpack_hook (self);
769790 }
770-
771- // If tensor_hold_helper is non-empty, some tensors may have been cleared by
772- // _clear_dataptr(). Iterate the top-level container tuple and restore any
773- // null impl from the corresponding entry in tensor_hold_helper.
774- // tensor_hold_helper is ordered by the DenseTensors found during deep
775- // traversal in set_container; for the common case (flat tuple of tensors)
776- // the k-th tensor in the tuple maps to tensor_hold_helper[k].
791+ // Re-attach any DenseTensor impls that were freed by _clear_dataptr().
792+ // tensor_hold_helper keeps the underlying allocations alive; walk the
793+ // container in the same DFS order as CollectDenseTensors and reinstall
794+ // impls for tensors whose impl() is currently null.
777795 if (!self->tensor_hold_helper .empty ()) {
778- Py_ssize_t size = PyTuple_Size (self->container );
779- PyObject* recovered_container = PyTuple_New (size);
780- Py_ssize_t holder_idx = 0 ;
781- for (Py_ssize_t i = 0 ; i < size; ++i) {
782- PyObject* item = PyTuple_GetItem (self->container , i);
783- if (item && PyCheckTensor (item)) {
784- TensorObject* tensor_obj = reinterpret_cast <TensorObject*>(item);
785- if (!tensor_obj->tensor .impl () &&
786- holder_idx <
787- static_cast <Py_ssize_t>(self->tensor_hold_helper .size ()) &&
788- self->tensor_hold_helper [holder_idx]) {
789- // Tensor was cleared by _clear_dataptr; restore impl from holder.
790- paddle::Tensor recovered;
791- recovered.set_impl (self->tensor_hold_helper [holder_idx]);
792- PyTuple_SET_ITEM (
793- recovered_container, i, paddle::pybind::ToPyObject (recovered));
794- ++holder_idx;
795- continue ;
796- }
797- ++holder_idx;
798- Py_INCREF (item);
799- PyTuple_SET_ITEM (recovered_container, i, item);
800- } else {
801- Py_INCREF (item);
802- PyTuple_SET_ITEM (recovered_container, i, item);
803- }
804- }
805- return recovered_container;
796+ RestoreDenseTensors (self->container , self->tensor_hold_helper );
806797 }
807-
808- // Fallback: return original container as-is.
809798 Py_INCREF (self->container );
810799 return self->container ;
811800 EAGER_CATCH_AND_THROW_RETURN_NULL
0 commit comments