Skip to content
Open
Show file tree
Hide file tree
Changes from all 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
135 changes: 135 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 @@ -56,6 +57,24 @@
]


_TENSOR_TYPE_DTYPES = {
'HalfTensor': 'float16',
'FloatTensor': 'float32',
'DoubleTensor': 'float64',
'Float8_e4m3fnTensor': 'float8_e4m3fn',
'Float8_e5m2Tensor': 'float8_e5m2',
'BFloat16Tensor': 'bfloat16',
'ByteTensor': 'uint8',
'CharTensor': 'int8',
'ShortTensor': 'int16',
'IntTensor': 'int32',
'LongTensor': 'int64',
'BoolTensor': 'bool',
'ComplexFloatTensor': 'complex64',
'ComplexDoubleTensor': 'complex128',
}


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


def _tensor_numel(input: Tensor) -> int:
"""
Returns the total number of elements in the tensor.

Args:
input (Tensor): The input tensor.

Returns:
int: The number of elements in ``input``.
"""
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.
"""
Returns the tensor dtype when ``dtype`` is not specified, otherwise casts
the tensor to the requested type.

Args:
input (Tensor): The input tensor.
dtype (DTypeLike|str|type|None, optional): The target tensor type or
data type. Qualified ``torch.*`` and ``paddle.*`` dtype or tensor
type strings are supported. When it is ``None``, returns a Paddle
dtype string. Default: ``None``.
non_blocking (bool, optional): Whether the conversion may occur
asynchronously. Default: ``False``.

Returns:
str|Tensor: A Paddle dtype string when ``dtype`` is ``None``;
otherwise, a tensor with the requested type.
"""
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:
return str(input.dtype)

device = None
if getattr(dtype, "__name__", None) in _TENSOR_TYPE_DTYPES:
# tensor factory classes, e.g. paddle.DoubleTensor
dtype = _TENSOR_TYPE_DTYPES[dtype.__name__]
device = "cpu"
elif isinstance(dtype, str):
dtype_string = dtype
tensor_type = dtype_string.rsplit(".", 1)[-1]
if not dtype_string.startswith(("torch.", "paddle.")):
raise ValueError(f"invalid type: {dtype_string!r}")
if tensor_type in _TENSOR_TYPE_DTYPES:
dtype = _TENSOR_TYPE_DTYPES[tensor_type]
device = (
"gpu"
if dtype_string.startswith(("torch.cuda.", "paddle.cuda."))
else "cpu"
)
elif tensor_type in _TENSOR_TYPE_DTYPES.values():
dtype = tensor_type
else:
raise ValueError(f"invalid type: {dtype_string!r}")

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,
)


@property
def _tensor_is_sparse(input: Tensor) -> bool:
Comment thread
Manfredss marked this conversation as resolved.
Comment thread
Manfredss marked this conversation as resolved.
"""
Whether the tensor uses the sparse COO layout.

Args:
input (Tensor): The input tensor.

Returns:
bool: ``True`` for a sparse COO tensor, otherwise ``False``.
"""
return input.is_sparse_coo()


def allclose(
input: Tensor,
other: Tensor,
Expand Down Expand Up @@ -1132,3 +1249,21 @@ def GetShapeOnDimInRange(shape, dim: int) -> int:
split_size_or_sections
)
return tuple(_C_ops.split(tensor, split_size_or_sections, dim))


# ``paddle.Tensor`` APIs routed to their ``paddle.compat`` implementations
_TENSOR_API_OVERRIDES = {
'allclose': allclose,
'equal': equal,
'slogdet': slogdet,
'sort': sort,
'split': split,
'min': min,
'max': max,
'unique': unique,
'median': median,
'nanmedian': nanmedian,
'numel': _tensor_numel,
'type': _tensor_type,
'is_sparse': _tensor_is_sparse,
}
89 changes: 50 additions & 39 deletions python/paddle/compat/api_dispatch.py
Original file line number Diff line number Diff line change
Expand Up @@ -51,29 +51,19 @@ def _caller_is_paddle_internal() -> bool:
return name == "paddle" or name.startswith("paddle.")


def dispatch_function(compat_fn: Any) -> Any:
"""Wrap a native ``paddle`` callable to route external callers to
``compat_fn`` while compat is enabled; paddle-internal callers and the
disabled state get the native callable. Installed only under
``enable_compat(level=2)``; ``disable_compat`` restores the originals,
so the default hot path is untouched."""

def decorator(native_fn: Any) -> Any:
@wraps(native_fn)
def dispatcher(*args: Any, **kwargs: Any) -> Any:
if (
len(_PADDLE_NAMESPACE_SAVED) > 0
and not _caller_is_paddle_internal()
):
return compat_fn(*args, **kwargs)
return native_fn(*args, **kwargs)
def dispatch_function(native_fn: Any, compat_fn: Any) -> Any:
"""Wrap a native ``paddle`` callable for caller-aware dispatch."""

dispatcher.__compat_fn__ = compat_fn
dispatcher.__native_fn__ = native_fn
dispatcher.__signature__ = inspect.signature(compat_fn)
return dispatcher
@wraps(native_fn)
def dispatcher(*args: Any, **kwargs: Any) -> Any:
if _caller_is_paddle_internal():
return native_fn(*args, **kwargs)
return compat_fn(*args, **kwargs)

return decorator
dispatcher.__compat_fn__ = compat_fn
dispatcher.__native_fn__ = native_fn
dispatcher.__signature__ = inspect.signature(compat_fn)
return dispatcher


def _iter_compat_modules() -> Generator[types.ModuleType, None, None]:
Expand Down Expand Up @@ -123,28 +113,49 @@ def __call__(cls, *args: Any, **kwargs: Any) -> Any:
return proxy


def dispatch_property(

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->func、func->property、property->property 三个分支可以共享这一个dispatch吗?

如果能共享,那后面三个分支也可以合并了

native_attr: Any,
compat_attr: Any,
) -> Any:
"""Route a Tensor API when either side uses the property protocol."""
compat_fn = (
compat_attr.fget if isinstance(compat_attr, property) else compat_attr
)

class _PropertyDispatcher:
def __get__(self, instance: Any, owner: type | None = None) -> Any:
if _caller_is_paddle_internal():
attr = native_attr
else:
attr = compat_attr
return attr.__get__(instance, owner)

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

这里能保证 attr 是 function 么?不然应该不能 .__get__

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.

不用一定是 function,只要是描述符就可以了吧

dispatch_property 只有在 _patch_tensor_methods 里才用到,native 这边来自 inspect.getattr_static(paddle.Tensor, name)compat 这边来自 paddle.compat.__all__(func)或 _TENSOR_API_OVERRIDES(func 或 property)。

现在只有两个 API 会走到这里:type(native 是 getset_descriptor,compat 是 func)、is_sparse(native 是 method_descriptor,compat 是 property)。这几个类型都实现了 __get__,所以我认为应该是没有问题的

@SigureMo SigureMo Aug 6, 2026

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

是 descriptor 还是 function 不重要,关键是能不能 bind object(__get__

算了,反正大概率也没人会在这两处注册 int 之类的 normal object,大概率也没啥问题,有问题再说吧

另外,其实 dispatch_property 已经能取代 dispatch_function 了吧?当然两者本质上有一点区别是,前者在 bind object 时候或者说 getattr 时候就已经确定 dispatch 的结果,而后者则是在实际函数调用时候才确定,但在实际调用过程中我觉得应该感知不到这点细微差异;还有一点差异就是之前可以用 paddle.Tensor.fn(obj) 而 descriptor 就不能 paddle.Tensor.des(obj) 了哈哈,但我觉得这倒没什么不至于有人这么去用,不过不重要了


dispatcher = _PropertyDispatcher()
dispatcher.__native_fn__ = native_attr
dispatcher.__compat_fn__ = compat_fn
dispatcher.__doc__ = compat_fn.__doc__
dispatcher.__name__ = compat_fn.__name__
dispatcher.__signature__ = inspect.signature(compat_fn)
return dispatcher


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
``x.max(dim=1)`` works torch-style for external callers (native for internal).
The dispatcher is patched directly like any paddle Tensor method: the
descriptor protocol forwards the tensor as the first positional argument,
which is exactly the compat function's ``input`` parameter.
"""
"""Route ``paddle.Tensor`` APIs to their root compat implementations."""
import paddle
import paddle.compat as compat_root

for attr_name in getattr(compat_root, "__all__", ()):
native_method = getattr(paddle.Tensor, attr_name, None)
if native_method is None:
for attr_name, compat_attr in compat_root._TENSOR_API_OVERRIDES.items():
native_attr = inspect.getattr_static(paddle.Tensor, attr_name, None)
if native_attr is None:
continue
compat_fn = getattr(compat_root, attr_name)
_PADDLE_NAMESPACE_SAVED[(paddle.Tensor, attr_name)] = native_method
setattr(
paddle.Tensor,
attr_name,
dispatch_function(compat_fn)(native_method),
)
_PADDLE_NAMESPACE_SAVED[(paddle.Tensor, attr_name)] = native_attr
if inspect.isdatadescriptor(native_attr) or isinstance(
compat_attr, property
):
dispatcher = dispatch_property(native_attr, compat_attr)
else:
dispatcher = dispatch_function(native_attr, compat_attr)
setattr(paddle.Tensor, attr_name, dispatcher)


def _apply_paddle_namespace_aliases() -> None:
Expand Down Expand Up @@ -177,7 +188,7 @@ def _apply_paddle_namespace_aliases() -> None:
setattr(
target_module,
attr_name,
dispatch_function(compat_attr)(current),
dispatch_function(current, compat_attr),
)
_patch_tensor_methods()

Expand Down
16 changes: 16 additions & 0 deletions python/paddle/compat/distributions/categorical.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,8 @@

from ..utils import _CompatClassMeta

__all__ = ["Categorical"]


class Categorical(distribution.Distribution, metaclass=_CompatClassMeta):
arg_constraints = {
Expand Down Expand Up @@ -66,6 +68,20 @@ def __init__(
distribution.Distribution.__init__(
self, batch_shape, validate_args=validate_args
)
if self._validate_args_enabled and paddle.in_dynamic_mode():
if probs is not None:
param_name = "probs"
valid = paddle.all(self.probs >= 0, axis=-1) & (
(self.probs.sum(-1) - 1).abs() < 1e-6
)
else:
param_name = "logits"
valid = constraint.real_vector.check(self.logits)
if not bool(valid.all()):
raise ValueError(
f'Expected parameter {param_name} of distribution '
'Categorical to satisfy its constraint'
)
Comment thread
Manfredss marked this conversation as resolved.

def expand(self, batch_shape, _instance=None):
new = (
Expand Down
Loading
Loading