Skip to content
Open
Show file tree
Hide file tree
Changes from 6 commits
Commits
Show all changes
28 commits
Select commit Hold shift + click to select a range
91b9144
align 3 Tensor api and PReLU
Manfredss Jul 27, 2026
3451263
Merge branch 'develop' of https://github.com/paddlepaddle/paddle into…
Manfredss Jul 27, 2026
d5154f3
fix PReLU, fix enable_compat restore after using guard
Manfredss Jul 27, 2026
4865ecc
add test coverage
Manfredss Jul 28, 2026
8aa81ef
fix per bot feedback
Manfredss Jul 28, 2026
7f2684a
fix typo
Manfredss Jul 28, 2026
7eda286
also fix paddle.distributions.categorical.Categorical
Manfredss Jul 28, 2026
ec4686b
fix
Manfredss Jul 29, 2026
cf76ec4
fix
Manfredss Jul 30, 2026
f668adb
Refine compat levels and guard state restoration
Manfredss Jul 30, 2026
111a665
Merge branch 'develop' of https://github.com/paddlepaddle/paddle into…
Manfredss Jul 30, 2026
b3bb47c
remove assertion
Manfredss Jul 30, 2026
2f3cc36
fix fleet tests failure
Manfredss Jul 31, 2026
ec1295c
Merge branch 'develop' of https://github.com/paddlepaddle/paddle into…
Manfredss Jul 31, 2026
f514e71
staged
Manfredss Jul 31, 2026
a721161
Merge branch 'develop' of https://github.com/paddlepaddle/paddle into…
Manfredss Aug 3, 2026
c64978f
reframe use_compat_guard
Manfredss Aug 3, 2026
826269a
remove unused methods
Manfredss Aug 3, 2026
a119dbe
fix counter; add dispatch_property
Manfredss Aug 3, 2026
5a5b7b3
fix
Manfredss Aug 3, 2026
219307c
fix tests
Manfredss Aug 3, 2026
fde90cf
Merge branch 'develop' of https://github.com/paddlepaddle/paddle into…
Manfredss Aug 4, 2026
b0217d4
Merge branch 'develop' of https://github.com/paddlepaddle/paddle into…
Manfredss Aug 5, 2026
82f8ebf
fix
Manfredss Aug 6, 2026
ce136ec
Merge branch 'develop' of https://github.com/paddlepaddle/paddle into…
Manfredss Aug 6, 2026
3eeb401
improve code per review suggestions
Manfredss Aug 6, 2026
43c95f1
fix
Manfredss Aug 7, 2026
8c035b1
refine
Manfredss Aug 7, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
89 changes: 89 additions & 0 deletions python/paddle/compat/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,7 @@
from collections.abc import Sequence

from paddle import Tensor
from paddle._typing import DTypeLike

__all__ = [
'allclose',
Expand All @@ -55,6 +56,26 @@
'seed',
]

_TENSOR_TYPE_NAMES = {
'float16': 'HalfTensor',
'float32': 'FloatTensor',
'float64': 'DoubleTensor',
'float8_e4m3fn': 'Float8_e4m3fnTensor',
'float8_e5m2': 'Float8_e5m2Tensor',
'bfloat16': 'BFloat16Tensor',
'uint8': 'ByteTensor',
'int8': 'CharTensor',
'int16': 'ShortTensor',
'int32': 'IntTensor',
'int64': 'LongTensor',
'bool': 'BoolTensor',
'complex64': 'ComplexFloatTensor',
'complex128': 'ComplexDoubleTensor',
}
_TENSOR_TYPE_DTYPES = {
tensor_type: dtype for dtype, tensor_type in _TENSOR_TYPE_NAMES.items()
}


def __getattr__(name):
if name == "paddle_triton":
Expand All @@ -66,6 +87,74 @@ def __getattr__(name):
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")


def _tensor_numel(input: Tensor) -> int:
return int(input.size)


def _tensor_type(
input: Tensor,
dtype: DTypeLike | str | type | None = None,
non_blocking: bool = False,
**kwargs: Any,
) -> str | Tensor:
Comment thread
Manfredss marked this conversation as resolved.
if "async" in kwargs:
non_blocking = kwargs.pop("async")
if kwargs:
key = next(iter(kwargs))
raise TypeError(f"type() got an unexpected keyword argument {key!r}")

if dtype is None:
dtype_name = str(input.dtype).removeprefix("paddle.")
tensor_type = _TENSOR_TYPE_NAMES[dtype_name]
prefix = "torch.cuda" if input.place.is_gpu_place() else "torch"

@zhwesky2010 zhwesky2010 Aug 2, 2026

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

这个返回 paddle.*

if input.is_sparse_coo():
prefix += ".sparse"
return f"{prefix}.{tensor_type}"

device = None
if isinstance(dtype, type) and dtype.__name__ in _TENSOR_TYPE_DTYPES:

@zhwesky2010 zhwesky2010 Aug 5, 2026

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

这里是否考虑了全部情况且能正确执行:

paddle.float32
np.float32
'float32'
'paddle.FloatTensor'
'torch.FloatTensor'

tensor_type = dtype.__name__
dtype = _TENSOR_TYPE_DTYPES[tensor_type]

@zhwesky2010 zhwesky2010 Aug 7, 2026

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

np.float32触发这个分支吗?看起来有点问题

直接判断np.dtype、paddle.dtype吧,鲁棒些。

这个isinstance(dtype, type) 比较奇怪

device = "cpu"
elif isinstance(dtype, str):
dtype_string = dtype
tensor_type = dtype_string.rsplit(".", 1)[-1]
if (

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

这里输入即可以是torch.xxx,也可以是paddle.xxx

not dtype_string.startswith("torch.")
or tensor_type not in _TENSOR_TYPE_DTYPES
):
raise ValueError(f"invalid type: {dtype_string!r}")
dtype = _TENSOR_TYPE_DTYPES[tensor_type]
device = "gpu" if dtype_string.startswith("torch.cuda.") else "cpu"

@zhwesky2010 zhwesky2010 Aug 2, 2026

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

paddle API进行兼容性设计:输入应同时支持torch.xxx和paddle.xxx,输出为paddle.xxx

这里好几个地方思路都不对


dtype_name = str(input.dtype).removeprefix("paddle.")

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

直接判断:str(input.dtype) == dtype 就可以了吧

dtype本身就是字符串

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

dtype 走到这里不一定是字符串好像,像t.type(paddle.float32)t.type(np.float32)t.type(np.dtype('float32')) 这些都不进上面两个分支,会直接传下来,而且即使是字符串,前面也已经把 'torch.DoubleTensor'、'paddle.float64' 这些统一成了 'float64',和带前缀的 str(input.dtype)('paddle.float64')比的话是不一致的 我还是想保留现在的写法

target_dtype_name = str(dtype).removeprefix("paddle.")
same_device = (
device is None
or (device == "cpu" and input.place.is_cpu_place())
or (device == "gpu" and input.place.is_gpu_place())
)
if dtype_name == target_dtype_name and same_device:
return input

return input.to(
device=device,
dtype=dtype,
blocking=not non_blocking,
)


def _tensor_is_sparse(input: Tensor) -> bool:
Comment thread
Manfredss marked this conversation as resolved.
Comment thread
Manfredss marked this conversation as resolved.
return input.is_sparse_coo()


_TENSOR_API_OVERRIDES = {

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

我建议这个是不是和上面的字典合并起来,统一成一个字典:

_TENSOR_API_OVERRIDES = (
    'allclose': allclose,
    'equal': equal
    'numel': _tensor_numel,
)

'numel': (_tensor_numel, False),
'type': (_tensor_type, False),
'is_sparse': (_tensor_is_sparse, True),
}


def allclose(
input: Tensor,
other: Tensor,
Expand Down
45 changes: 45 additions & 0 deletions python/paddle/compat/api_dispatch.py
Original file line number Diff line number Diff line change
Expand Up @@ -123,6 +123,37 @@ def __call__(cls, *args: Any, **kwargs: Any) -> Any:
return proxy


class _TensorCompatDescriptor:

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

这个与其他tensor method这里为何会有这么多特殊之处?除了名字不同。

这一块的设计比较冗余,优化下设计,与其他tensor method合并处理。尽可能代码复用并减少行数。

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

不太理解要加这个东西解决什么问题,如果只是一个property的问题,单独整一个 dispatch_property 这个函数,其他继续使用dispatch_function

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

这个问题点改了吗?上次沟通的意思理解了没?看起来和沟通的不是一回事

disptch_function + disptch_property

"""Caller-aware adapter for Tensor method/property shape differences."""

def __init__(
self,
native_attr: Any,
compat_fn: Any,
as_property: bool,
) -> None:
self.__native_fn__ = native_attr
self.__compat_fn__ = compat_fn
self._as_property = as_property
self.__doc__ = compat_fn.__doc__
self.__name__ = compat_fn.__name__
self.__signature__ = inspect.signature(compat_fn)

def __get__(self, instance: Any, owner: type | None = None) -> Any:
if instance is None:
if _caller_is_paddle_internal():
return self.__native_fn__
return self if self._as_property else self.__compat_fn__
if (
len(_PADDLE_NAMESPACE_SAVED) > 0
and not _caller_is_paddle_internal()
):
if self._as_property:
return self.__compat_fn__(instance)
return self.__compat_fn__.__get__(instance, owner)
return self.__native_fn__.__get__(instance, owner)


def _patch_tensor_methods() -> None:
"""Route ``paddle.Tensor.<m>`` to the compat function for the root compat APIs
that torch also exposes as Tensor methods (max/min/sort/split/unique/...), so
Expand All @@ -146,6 +177,20 @@ def _patch_tensor_methods() -> None:
dispatch_function(compat_fn)(native_method),
)

for attr_name, (

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

这个 _TENSOR_API_OVERRIDES 和 __all__先合并为字典,然后再统一进行patch

compat_fn,
as_property,
) in compat_root._TENSOR_API_OVERRIDES.items():
native_attr = inspect.getattr_static(paddle.Tensor, attr_name, None)
if native_attr is None:
continue
_PADDLE_NAMESPACE_SAVED[(paddle.Tensor, attr_name)] = native_attr
setattr(
paddle.Tensor,
attr_name,
_TensorCompatDescriptor(native_attr, compat_fn, as_property),
)


def _apply_paddle_namespace_aliases() -> None:
"""Install caller-aware dispatchers/proxies for every public ``paddle.compat.*``
Expand Down
17 changes: 17 additions & 0 deletions python/paddle/compat/nn/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,7 @@
__all__ = [
'Unfold',
'Linear',
'PReLU',
'Softmax',
'AvgPool1D',
'AvgPool2D',
Expand Down Expand Up @@ -662,6 +663,22 @@ def reset_parameters(self) -> None:
nn.init.uniform_(self.bias, -bound, bound)


class PReLU(nn.PReLU, metaclass=_CompatClassMeta):
def __init__(
self,
num_parameters: int = 1,
init: float = 0.25,
device: PlaceLike | None = None,

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

这个做一个重载设计吧,新增一个装饰器,判断传入的第3个位置参数是ParamAttrLike还是PlaceLike类型进行分发

非必要不添加compat函数

dtype: DTypeLike | None = None,
) -> None:
super().__init__(
num_parameters=num_parameters,
init=init,
device=device,
dtype=dtype,
)


class Softmax(nn.Layer, metaclass=_CompatClassMeta):
r"""
Softmax Activation.
Expand Down
4 changes: 3 additions & 1 deletion python/paddle/compat/proxy.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@
from typing import TYPE_CHECKING, Any, Literal

from .api_dispatch import (
_PADDLE_NAMESPACE_SAVED,

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

只是加一个level,这里为啥要改动上百行,5行内能解决问题吗?

_apply_paddle_namespace_aliases,
_iter_compat_modules,
_restore_paddle_namespace_aliases,
Expand Down Expand Up @@ -604,6 +605,7 @@ def use_compat_guard(
already_has_torch_proxy = TORCH_PROXY_FINDER in sys.meta_path
original_local_enabled_scope = set(TORCH_PROXY_FINDER._local_enabled_scope)
original_globally_enabled = TORCH_PROXY_FINDER._globally_enabled
original_level = 2 if _PADDLE_NAMESPACE_SAVED else 1
if enable == already_has_torch_proxy and (
(original_globally_enabled and scope is None)
or (original_local_enabled_scope == (scope or set()))
Expand All @@ -625,7 +627,7 @@ def use_compat_guard(
try:
yield
finally:
enable_compat(scope=None, silent=True)
enable_compat(scope=None, silent=True, level=original_level)
TORCH_PROXY_FINDER._local_enabled_scope = (
original_local_enabled_scope
)
Expand Down
3 changes: 3 additions & 0 deletions python/paddle/nn/layer/activation.py
Original file line number Diff line number Diff line change
Expand Up @@ -586,6 +586,9 @@ def __init__(
device: PlaceLike | None = None,
dtype: DTypeLike | None = None,
) -> None:
if not isinstance(data_format, str):

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

自定义装饰器 兼容两套签名,参考skill的具体要求

device, dtype = weight_attr, data_format
weight_attr, data_format = None, "NCHW"
super().__init__()
self._num_parameters = num_parameters
self._init = init
Expand Down
Loading
Loading