torch

API名称

是否支持

限制与说明

torch.is_tensor

是

  

torch.is_storage

是

  

torch.is_complex

是

支持判断,但当前硬件限制不支持复数。

torch.is_conj

是

  

torch.is_floating_point

是

  

torch.is_nonzero

是

  

torch.set_default_dtype

是

  

torch.get_default_dtype

是

  

torch.set_default_tensor_type

否

  

torch.numel

是

  

torch.set_printoptions

是

  

torch.set_flush_denormal

是

  

torch.tensor

是

  

torch.sparse_coo_tensor

否

  

torch.asarray

是

  

torch.as_tensor

是

  

torch.as_strided

是

  

torch.from_numpy

是

  

torch.frombuffer

是

  

torch.zeros

是

  

torch.zeros_like

是

  

torch.ones

是

  

torch.ones_like

是

  

torch.arange

是

  

torch.range

是

  

torch.linspace

是

  

torch.logspace

是

  

torch.eye

是

  

torch.empty

是

  

torch.empty_like

是

  

torch.empty_strided

是

  

torch.full

是

  

torch.full_like

是

  

torch.quantize_per_tensor

是

  

torch.quantize_per_channel

是

  

torch.dequantize(tensor) ->Tensor

否

  

torch.dequantize(tensors) ->sequence of Tensors

否

  

torch.complex

否

  

torch.polar

否

  

torch.heaviside

否

  

torch.adjoint

否

  

torch.argwhere

否

  

torch.cat

是

  

torch.concat

是

不支持float64,不支持8D输入。

torch.conj

否

  

torch.chunk

是

  

torch.dsplit

是

只支持float16,float32,float64。

torch.column_stack

是

  

torch.dstack

是

  

torch.gather

是

  

torch.hsplit

是

  

torch.hstack

是

  

torch.index_add

是

  

torch.index_select

是

  

torch.masked_select

是

  

torch.movedim

是

  

torch.moveaxis

是

  

torch.narrow

是

  

torch.nonzero

是

  

torch.permute

是

  

torch.reshape

是

  

torch.row_stack

是

  

torch.select

是

  

torch.scatter

是

  

torch.diagonal_scatter

是

  

torch.select_scatter

是

  

torch.slice_scatter

是

  

torch.scatter_add

是

  

torch.scatter_reduce

否

  

torch.split

是

  

torch.squeeze

是

  

torch.stack

是

  

torch.swapaxes

是

  

torch.swapdims

是

  

torch.t

是

  

torch.take

是

  

torch.take_along_dim

是

不支持float64与int。

torch.tensor_split

是

  

torch.tile

是

  

torch.transpose

是

  

torch.unbind

是

  

torch.unsqueeze

是

  

torch.vsplit

是

只支持float16,float32,float64。

torch.vstack

是

  

torch.where(condition, x, y) ->Tensor

是

  

torch.where(condition) ->tuple of LongTensor

是

  

torch.Generator

否

  

torch.Generator.get_state

否

  

torch.Generator.initial_seed

否

  

torch.Generator.manual_seed

否

  

torch.Generator.seed

否

  

torch.Generator.set_state

否

  

torch.seed

是

  

torch.manual_seed

是

  

torch.initial_seed

是

  

torch.get_rng_state

是

  

torch.set_rng_state

是

  

torch.bernoulli

是

  

torch.multinomial

是

  

torch.normal(mean, std, *, generator=None, out=None) ->Tensor

否

  

torch.normal(mean=0.0, std, *, out=None) ->Tensor

否

  

torch.normal(mean, std=1.0, *, out=None) ->Tensor

否

  

torch.normal(mean, std, size, *, out=None) ->Tensor

否

  

torch.poisson

否

可以走CPU实现。

torch.rand

是

  

torch.rand_like

是

不支持int64。

torch.randint

是

  

torch.randint_like

是

  

torch.randn

是

  

torch.randn_like

是

  

torch.randperm

是

  

torch.quasirandom.SobolEngine

是

  

torch.quasirandom.SobolEngine.draw

是

  

torch.quasirandom.SobolEngine.draw_base2

是

  

torch.quasirandom.SobolEngine.fast_forward

是

  

torch.quasirandom.SobolEngine.reset

是

  

torch.save

是

  

torch.load

是

  

torch.get_num_threads

是

仅支持在CPU上运行。

torch.set_num_threads

是

仅支持在CPU上运行。

torch.get_num_interop_threads

是

  

torch.set_num_interop_threads

是

仅支持在CPU上运行。

torch.no_grad

是

  

torch.enable_grad

是

  

torch.set_grad_enabled

是

  

torch.is_grad_enabled

是

  

torch.inference_mode

是

  

torch.is_inference_mode_enabled

是

  

torch.abs

是

  

torch.absolute

是

  

torch.acos

是

不支持int32。

torch.arccos

是

不支持int64。

torch.acosh

是

不支持int64。

torch.arccosh

是

不支持int64。

torch.add

是

  

torch.addcdiv

是

支持fp16,fp32,int64,bool。

在int64类型不支持三个tensor同时广播。

torch.addcmul

是

支持fp16,fp32,fp64,int8,int32,int64,uint8,bool,bf16。

在int8、uint8、int64、fp64类型不支持三个tensor同时广播。

torch.angle

是

可以走CPU实现。

torch.asin

是

  

torch.arcsin

是

不支持int64。

torch.asinh

是

  

torch.arcsinh

是

不支持int64。

torch.atan

是

  

torch.arctan

是

不支持int64。

torch.atanh

是

不支持int64。

torch.arctanh

是

  

torch.atan2

是

  

torch.arctan2

是

  

torch.bitwise_not

是

  

torch.bitwise_and

是

  

torch.bitwise_or

是

  

torch.bitwise_xor

是

  

torch.bitwise_left_shift

否

  

torch.bitwise_right_shift

否

  

torch.ceil

是

  

torch.clamp

是

  

torch.clip

是

  

torch.conj_physical

否

  

torch.copysign

是

  

torch.cos

是

  

torch.cosh

是

  

torch.deg2rad

是

  

torch.div

是

  

torch.divide

否

  

torch.digamma

否

可以走CPU实现。

torch.erf

是

  

torch.erfc

是

  

torch.erfinv

是

  

torch.exp

是

  

torch.exp2

是

  

torch.expm1

是

  

torch.fake_quantize_per_channel_affine

否

  

torch.fake_quantize_per_tensor_affine

否

  

torch.fix

是

  

torch.float_power

是

不支持float64。

torch.floor

是

  

torch.floor_divide

是

  

torch.fmod

是

  

torch.frac

是

  

torch.frexp

否

  

torch.gradient

是

不支持8D输入。

torch.imag

否

  

torch.ldexp

是

在int64带out场景下,out场景也必须是int64类型。

torch.lerp

是

  

torch.lgamma

否

可以走CPU实现。

torch.log

是

  

torch.log10

是

  

torch.log1p

是

  

torch.log2

是

  

torch.logaddexp

否

不支持double数据类型。

torch.logaddexp2

否

不支持double数据类型。

torch.logical_and

是

  

torch.logical_not

是

  

torch.logical_or

是

  

torch.logical_xor

是

可以走CPU实现。

torch.logit

否

可以走CPU实现。

torch.hypot

否

  

torch.i0

否

  

torch.igamma

否

  

torch.igammac

否

  

torch.mul

是

  

torch.multiply

是

  

torch.mvlgamma

否

可以走CPU实现。

torch.nan_to_num

否

  

torch.neg

是

  

torch.negative

是

  

torch.nextafter

否

  

torch.polygamma

否

  

torch.positive

是

  

torch.pow(input, exponent, *, out=None) ->Tensor

是

不支持int64。

Int类型使用的是浮点数方案,可能存在精度误差。

torch.pow(self, exponent, *, out=None) ->Tensor

是

不支持int64。

Int类型使用的是浮点数方案,可能存在精度误差。

torch.quantized_batch_norm

否

  

torch.quantized_max_pool1d

否

  

torch.quantized_max_pool2d

否

  

torch.rad2deg

是

不支持int64。

torch.real

是

  

torch.reciprocal

是

  

torch.remainder

否

  

torch.round

是

  

torch.rsqrt

是

  

torch.sigmoid

是

  

torch.sign

是

  

torch.sgn

否

  

torch.signbit

否

  

torch.sin

是

  

torch.sinc

否

  

torch.sinh

是

  

torch.sqrt

是

  

torch.square

是

  

torch.sub

是

  

torch.subtract

是

  

torch.tan

是

  

torch.tanh

是

  

torch.true_divide

是

  

torch.trunc

是

  

torch.xlogy

否

  

torch.argmax(input) ->LongTensor

是

  

torch.argmax(input, dim, keepdim=False) ->LongTensor

是

  

torch.argmin

是

  

torch.amax

是

  

torch.amin

是

  

torch.aminmax

否

  

torch.all(input) ->Tensor

是

  

torch.all(input, dim, keepdim=False, *, out=None) ->Tensor

是

  

torch.any(input) ->Tensor

是

  

torch.any(input, dim, keepdim=False, *, out=None) ->Tensor

是

  

torch.max(input) ->Tensor

是

  

torch.max(input, dim, keepdim=False, *, out=None)

是

  

torch.max(input, other, *, out=None) ->Tensor

是

  

torch.min(input) ->Tensor

是

  

torch.min(input, dim, keepdim=False, *, out=None)

是

  

torch.min(input, other, *, out=None) ->Tensor

是

  

torch.dist

是

  

torch.logsumexp

是

  

torch.mean(input, *, dtype=None) ->Tensor

是

  

torch.mean(input, dim, keepdim=False, *, dtype=None, out=None) ->Tensor

是

  

torch.nanmean

否

  

torch.median(input) ->Tensor

是

  

torch.median(input, dim=- 1, keepdim=False, *, out=None)

是

  

torch.nanmedian(input) ->Tensor

否

可以走CPU实现。

torch.nanmedian(input, dim=- 1, keepdim=False, *, out=None)

否

可以走CPU实现。

torch.mode

否

可以走CPU实现。

torch.norm

是

  

torch.nansum(input, *, dtype=None) ->Tensor

否

可以走CPU实现。

torch.nansum(input, dim, keepdim=False, *, dtype=None) ->Tensor

否

可以走CPU实现。

torch.prod(input, *, dtype=None) ->Tensor

是

  

torch.prod(input, dim, keepdim=False, *, dtype=None) ->Tensor

是

  

torch.quantile

是

  

torch.nanquantile

否

  

torch.std(input, dim, unbiased, keepdim=False, *, out=None) ->Tensor

是

如果输入tensor元素值相同,会产生精度误差。

torch.std(input, unbiased) ->Tensor

是

如果输入tensor元素值相同,会产生精度误差。

torch.std_mean(input, dim, unbiased, keepdim=False, *, out=None)

是

只支持float16,float32。

torch.std_mean(input, unbiased)

是

只支持float16,float32。

torch.sum(input, *, dtype=None) ->Tensor

是

  

torch.sum(input, dim, keepdim=False, *, dtype=None) ->Tensor

是

  

torch.unique

是

  

torch.unique_consecutive

是

传参时必须使用关键字,否则精度不达标。return_inverse=return_inverse,return_counts=return_counts,dim=dim。

torch.var(input, dim, unbiased, keepdim=False, *, out=None) ->Tensor

是

  

torch.var(input, unbiased) ->Tensor

是

  

torch.var_mean(input, dim, unbiased, keepdim=False, *, out=None)

是

  

torch.var_mean(input, unbiased)

是

  

torch.count_nonzero

是

  

torch.allclose

是

  

torch.argsort

是

  

torch.eq

是

  

torch.equal

是

  

torch.ge

是

  

torch.greater_equal

是

  

torch.gt

是

  

torch.greater

是

  

torch.isclose

是

  

torch.isfinite

是

  

torch.isin

否

  

torch.isinf

是

  

torch.isposinf

是

  

torch.isneginf

是

  

torch.isnan

Atlas 训练系列产品:是

Atlas A2 训练系列产品:否

  

torch.isreal

是

  

torch.kthvalue

否

  

torch.le

否

  

torch.less_equal

是

  

torch.lt

否

  

torch.less

是

  

torch.maximum

是

  

torch.minimum

是

  

torch.fmax

否

  

torch.fmin

是

可以走CPU实现。

torch.ne

否

  

torch.not_equal

是

  

torch.sort

否

  

torch.topk

否

  

torch.msort

否

  

torch.stft

否

  

torch.istft

否

  

torch.bartlett_window

否

  

torch.blackman_window

否

  

torch.hamming_window

否

  

torch.hann_window

否

  

torch.kaiser_window

否

  

torch.atleast_1d

是

  

torch.atleast_2d

是

  

torch.atleast_3d

是

  

torch.bincount

是

  

torch.block_diag

是

  

torch.broadcast_tensors

否

  

torch.broadcast_to

是

  

torch.broadcast_shapes

是

  

torch.bucketize

否

可以走CPU实现。

torch.cartesian_prod

是

  

torch.cdist

是

  

torch.clone

是

  

torch.combinations(input, r=2, with_replacement=False) ->seq

是

r不能大于8。

torch.corrcoef

是

不支持float64。

torch.cov

是

不支持bool。

torch.cross

是

两个输入的shape要保持一致。

torch.cummax

是

  

torch.cummin

是

  

torch.cumprod

是

  

torch.cumsum

是

  

torch.diag

是

仅支持diagonal=0场景。

torch.diag_embed

是

不支持复数。

torch.diagflat

是

  

torch.diagonal

是

  

torch.diff

是

  

torch.einsum

否

  

torch.flatten

是

  

torch.flip

是

  

torch.fliplr

是

  

torch.flipud

是

  

torch.kron

是

不支持5维度及以上输入。

torch.rot90

是

  

torch.gcd

是

可以走CPU实现。

torch.histc

否

  

torch.histogram

否

  

torch.histogramdd

否

  

torch.meshgrid

是

  

torch.lcm

否

  

torch.logcumsumexp

否

  

torch.ravel

是

  

torch.renorm

是

支持fp16,fp32

属性max_norm仅支持非负值

torch.repeat_interleave(input, repeats, dim=None, *, output_size=None) ->Tensor

是

  

torch.repeat_interleave(repeats, *, output_size=None) ->Tensor

是

  

torch.roll

是

  

torch.searchsorted

是

  

torch.tensordot

是

  

torch.trace

是

支持fp16,fp32,fp64,int8,int16,int32,int64,uint8,complex64,complex128

torch.tril

是

  

torch.tril_indices

是

  

torch.triu

是

  

torch.triu_indices

是

  

torch.vander

否

  

torch.view_as_real

否

  

torch.view_as_complex

否

  

torch.resolve_conj

否

  

torch.resolve_neg

否

  

torch.addbmm

是

  

torch.addmm

是

  

torch.addmv

是

  

torch.addr

是

  

torch.baddbmm

是

  

torch.bmm

是

  

torch.chain_matmul

是

  

torch.cholesky

否

  

torch.cholesky_inverse

否

  

torch.cholesky_solve

否

  

torch.dot

是

  

torch.eig

否

  

torch.geqrf

否

  

torch.ger

是

  

torch.inner

是

不支持int64,out场景没适配。

torch.inverse

是

  

torch.det

否

  

torch.logdet

否

  

torch.slogdet

是

  

torch.lstsq

否

  

torch.lu

否

  

torch.lu_solve

否

  

torch.lu_unpack

否

  

torch.matmul

是

  

torch.matrix_power

是

  

torch.matrix_rank

是

不支持symmetric=True,out场景没适配。

torch.matrix_exp

否

  

torch.mm

是

  

torch.mv

是

  

torch.orgqr

否

  

torch.ormqr

否

  

torch.outer

是

  

torch.pinverse

是

  

torch.qr

是

  

torch.solve

否

  

torch.svd

是

  

torch.svd_lowrank

否

可以走CPU实现。

torch.pca_lowrank

否

可以走CPU实现。

torch.symeig

是

  

torch.lobpcg

否

  

torch.trapz

是

  

torch.trapezoid

是

  

torch.cumulative_trapezoid

是

只支持float16,float32。

torch.triangular_solve

是

  

torch.vdot

否

可以走CPU实现。

torch.compiled_with_cxx11_abi

是

  

torch.result_type

是

  

torch.can_cast

是

  

torch.promote_types

是

  

torch.use_deterministic_algorithms

是

  

torch.are_deterministic_algorithms_enabled

是

  

torch.is_deterministic_algorithms_warn_only_enabled

否

  

torch.set_deterministic_debug_mode

是

  

torch.get_deterministic_debug_mode

是

  

torch.set_warn_always

是

  

torch.is_warn_always_enabled

是

  

torch._assert

是