如何为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
相关产品推荐
相关产品推荐

