You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何为jaxtyping中带参数的类型创建Python类型别名

自定义jaxtyping类型别名的Python实现

jaxtyping支持以字符串作为参数(而非类型)的类型注解,示例写法如下:

Float[Array, "dim1 dim2"]

我想定义一个合并Float和Array的类型别名,这样就能用MyOwnType["dim1 dim2"]替代上面的写法。但我发现这里的泛型参数"dim1 dim2"是字符串实例而非类型,没法用TypeAlias或TypeVar来实现。

有没有简洁的、符合Python风格的解法?


尝试过的无效代码

我试过下面的实现,但并不奏效:

class _Singleton:
    def __getitem__(self, shape: str) -> Float:
        return Float[Array, shape]

MyOwnType = _Singleton()

当把MyOwnType["dim1 dim2"]用作函数参数注解时,mypy会抛出错误:Variable "MyOwnType" is not valid as a type


最终可行方案

基于@chepner的回答,最终的有效代码如下:

class MyOwnType(Generic[Shape]):
    def __class_getitem__(cls, shape: str) -> Float:
        return Float[Array, shape]

内容的提问来源于stack exchange,提问作者Padix Key

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.21 00:40:06