PyTorch torch Reference Manual

The PyTorch package contains data structures for multidimensional tensors and defines mathematical operations performed on these tensors. In addition, it provides many utility tools for efficiently serializing tensors and arbitrary types of data, along with other useful utilities.

It also has a CUDA version that allows you to run tensor computations on NVIDIA GPUs with compute capability >= 3.0.

PyTorch torch API Manual

Tensor Type Checks

Function Description
torch.is_tensor(obj) Checkobjwhether it is a PyTorch tensor.
torch.is_storage(obj) Checkobjwhether it is a PyTorch storage object.
torch.is_complex(input) Checkinputwhether the data type is a complex data type.
torch.is_conj(input) Checkinputwhether it is a conjugate tensor.
torch.is_floating_point(input) Checkinputwhether the data type is a floating-point data type.
torch.is_nonzero(input) Checkinputwhether it is a non-zero single-element tensor.
torch.set_default_dtype(d) Set the default floating-point data type tod。
torch.get_default_dtype() Get the current default floating-pointtorch.dtype。
torch.set_default_device(device) Set the defaulttorch.Tensorallocated device todevice。
torch.get_default_device() Get the defaulttorch.Tensorallocated device.
torch.set_default_tensor_type(tensor_type) Set the default tensor type totensor_type。
torch.numel(input) Returninputthe total number of elements in the tensor.
torch.set_printoptions(...) Set tensor printing options.

Tensor Creation

Function Description
torch.tensor(data, dtype, device, requires_grad) Create a tensor from data, copying the data, with no autograd history.
torch.as_tensor(data, dtype, device) Convert data to a tensor, sharing data and preserving autograd history.
torch.asarray(data, dtype, device) Convert data to a tensor array.
torch.from_numpy(ndarray) Create a tensor from a NumPy array (shared memory).
torch.from_dlpack(ext_tensor) Create a PyTorch tensor from a dlpack tensor.
torch.frombuffer(buffer, dtype, count, offset) Create a 1-D tensor from a buffer.
torch.zeros(*size, dtype, device, requires_grad) Create a tensor of all zeros.
torch.zeros_like(input, dtype, device, requires_grad) Create a tensor of all zeros with the same shape as the input.
torch.ones(*size, dtype, device, requires_grad) Create a tensor of all ones.
torch.ones_like(input, dtype, device, requires_grad) Create a tensor of all ones with the same shape as the input.
torch.empty(*size, dtype, device, requires_grad) Create an uninitialized tensor.
torch.empty_like(input, dtype, device, requires_grad) Create an uninitialized tensor with the same shape as the input.
torch.empty_strided(size, stride, dtype, device) Create an uninitialized tensor with specified strides.
torch.arange(start, end, step, dtype, device, requires_grad) Create an arithmetic sequence tensor.
torch.range(start, end, step, dtype, device, requires_grad) Create an arithmetic sequence tensor that includes the end value.
torch.linspace(start, end, steps, dtype, device, requires_grad) Create an evenly spaced sequence tensor.
torch.logspace(start, end, steps, base, dtype, device, requires_grad) Create a logarithmically spaced sequence tensor.
torch.eye(n, m, dtype, device, requires_grad) Create an identity matrix.
torch.full(size, fill_value, dtype, device, requires_grad) Create a tensor filled with a specified value.
torch.full_like(input, fill_value, dtype, device, requires_grad) Create a tensor filled with a value and having the same shape as the input.
torch.rand(*size, dtype, device, requires_grad) Create a uniformly distributed random tensor (range [0, 1)).
torch.rand_like(input, dtype, device, requires_grad) Create a uniformly distributed random tensor with the same shape as the input.
torch.randn(*size, dtype, device, requires_grad) Create a standard normal distribution random tensor.
torch.randn_like(input, dtype, device, requires_grad) Create a standard normal distribution random tensor with the same shape as the input.
torch.randint(low, high, size, dtype, device, requires_grad) Create an integer random tensor.
torch.randint_like(input, low, high, dtype, device, requires_grad) Create an integer random tensor with the same shape as the input.
torch.randperm(n, dtype, device, requires_grad) Create a random permutation of 0 to n-1.
torch.sparse_coo_tensor(indices, values, size, dtype, device, requires_grad) At the specifiedindicesposition, construct a sparse COO tensor.
torch.sparse_csr_tensor(crow_indices, col_indices, values, size, dtype, device) Construct a sparse CSR tensor.
torch.sparse_csc_tensor(ccol_indices, row_indices, values, size, dtype, device) Construct a sparse CSC tensor.
torch.quantize_per_tensor(input, scale, zero_point, dtype) Create a quantized tensor (per-tensor).
torch.quantize_per_channel(input, scales, zero_points, axis, dtype) Create a quantized tensor (per-channel).
torch.dequantize(input) Dequantize a tensor.
torch.complex(real, imag) Create a complex tensor from real and imaginary parts.
torch.polar(abs, angle) Create a complex tensor from polar coordinates.
torch.heaviside(input, values) Compute the Heaviside step function.

Indexing, Slicing, Joining, Mutating Operations

Function Description
torch.cat(tensors, dim, out) Concatenate tensors along a specified dimension.
torch.concat(tensors, dim, out) Concatenate tensors along a specified dimension (same as cat).
torch.concatenate(tensors, dim, out) Concatenate tensors along a specified dimension (same as cat).
torch.stack(tensors, dim, out) Stack tensors along a new dimension.
torch.split(tensor, split_size, dim) Split a tensor along a specified dimension.
torch.chunk(tensor, chunks, dim) Chunk a tensor along a specified dimension.
torch.reshape(input, shape) Change the shape of a tensor.
torch.transpose(input, dim0, dim1) Swap two dimensions of a tensor.
torch.t(input) Transpose a 2-D tensor.
torch.squeeze(input, dim) Remove a dimension of size 1.
torch.unsqueeze(input, dim) Insert a dimension of size 1 at a specified position.
torch.permute(input, dims) Rearrange the dimensions of a tensor.
torch.movedim(input, source, destination) Move a dimension of the tensor to a new position.
torch.moveaxis(input, source, destination) Move an axis of the tensor to a new position.
torch.narrow(input, dim, start, length) Return a slice of a tensor.
torch.narrow_copy(input, dim, start, length) Return a copy of a slice of a tensor.
torch.select(input, dim, index) Select slices corresponding to indices along a specified dimension.
torch.slice_scatter(input, src, dim, start, end) Scatter src into the slices of input.
torch.select_scatter(input, src, dim, index) Scatter src into the specified index positions.
torch.diagonal_scatter(input, src, offset, dim1, dim2) Scatter values into diagonal positions.
torch.expand(input, size) Expand the size of a tensor (copy view).
torch.expand_as(input, other) Expand the tensor to the same size as other.
torch.masked_select(input, mask) Select elements according to a boolean mask.
torch.index_select(input, dim, index) Select elements corresponding to indices along a specified dimension.
torch.gather(input, dim, index, sparse_grad) Gather elements at specified indices along a specified dimension.
torch.scatter(input, dim, index, src, reduce) willsrcscatter the values of ... toinputthe specified positions of ...
torch.scatter_add(input, dim, index, src) Add the values of src to the specified positions.
torch.scatter_reduce(input, dim, index, src, reduce, include_self) Aggregate the values of src to the specified positions in the specified manner.
torch.index_add(input, dim, index, source, alpha) Add source to the positions specified by index.
torch.index_copy(input, dim, index, source) Copy source to the positions specified by index.
torch.index_reduce(input, dim, index, source, reduce, include_self) Aggregate source to the positions specified by index in the specified manner.
torch.take(input, index) Get the element at the given index position.
torch.take_along_dim(input, indices, dim) Get the elements at index positions along a specified dimension.
torch.nonzero(input) Return the indices of non-zero elements.
torch.argwhere(input) Return the indices of elements that satisfy the condition.
torch.where(condition, input, other) Return elements according to a condition.
torch.unbind(tensor, dim) Split into a tuple along a specified dimension.
torch.split_with_sizes(tensor, split_sizes, dim) Split a tensor by sizes.
torch.tensor_split(tensor, indices_or_sections, dim) Split a tensor by indices or number of segments.
torch.hsplit(tensor, indices_or_sections) Split a tensor horizontally.
torch.vsplit(tensor, indices_or_sections) Split a tensor vertically.
torch.dsplit(tensor, indices_or_sections) Split a tensor depthwise.
torch.hstack(tensors, dim, out) Stack tensors horizontally.
torch.vstack(tensors, out) Stack tensors vertically.
torch.dstack(tensors, out) Stack tensors depthwise.
torch.column_stack(tensors, out) Stack tensors column-wise.
torch.row_stack(tensors, out) Stack tensors row-wise (same as vstack).
torch.tile(input, dims) Repeat a tensor multiple times.
torch.repeat_interleave(input, repeats, dim) Repeat elements along a specified dimension.
torch.flip(input, dims) Flip a tensor along a specified dimension.
torch.fliplr(input) Flip a tensor left-right.
torch.flipud(input) Flip a tensor up-down.
torch.rot90(input, k, dims) Rotate a tensor by 90 degrees.
torch.linalg.matrix_transpose(input) Matrix transpose.
torch.adjoint(input) Return the adjoint of a tensor.
torch.resolve_conj(input) Resolve the conjugate tensor.
torch.resolve_neg(input) Resolve the negative tensor.
torch.view_as_real(input) Treat a complex tensor as a real tensor.
torch.view_as_complex(input) Treat a real tensor as a complex tensor.
torch.unravel_index(indices, shape) Convert flattened indices to multi-dimensional indices.

Random Number Generation

Function Description
torch.manual_seed(seed) Set the random seed (CPU).
torch.seed() Set the random seed and return the new seed value.
torch.initial_seed() Return the current random seed.
torch.get_rng_state() Return the random number generator state.
torch.set_rng_state(state) Set the random number generator state.
torch.rand(*size, dtype, device, requires_grad) Create a uniformly distributed random tensor (range [0, 1)).
torch.rand_like(input, dtype, device, requires_grad) Create a uniformly distributed random tensor with the same shape as the input.
torch.randn(*size, dtype, device, requires_grad) Create a standard normal distribution random tensor.
torch.randn_like(input, dtype, device, requires_grad) Create a standard normal distribution random tensor with the same shape as the input.
torch.randint(low, high, size, dtype, device, requires_grad) Create an integer random tensor.
torch.randint_like(input, low, high, dtype, device, requires_grad) Create an integer random tensor with the same shape as the input.
torch.randperm(n, dtype, device, requires_grad) Create a random permutation of 0 to n-1.
torch.bernoulli(input, *, generator) Generate random numbers from a Bernoulli distribution.
torch.multinomial(input, num_samples, replacement, generator) Multinomial sampling.
torch.normal(mean, std, out) Generate random numbers from a normal distribution.
torch.poisson(input, generator) Generate random numbers from a Poisson distribution.

Serialization

Function Description
torch.save(obj, f, pickle_module, pickle_protocol) Save an object to a file.
torch.load(f, map_location, pickle_module, weights_only) Load an object from a file.

Gradient Control

Function Description
torch.no_grad() Context manager that disables gradient computation.
torch.enable_grad() Context manager that enables gradient computation.
torch.set_grad_enabled(grad) Set whether gradient computation is enabled.
torch.is_grad_enabled() Check whether gradient computation is enabled.
torch.inference_mode() Context manager for inference mode (disables gradients and autograd).
torch.is_inference_mode_enabled() Check whether inference mode is enabled.

Mathematical operations - pointwise operations

Function Description
torch.abs(input, out) Element-wise absolute value.
torch.absolute(input, out) Element-wise absolute value (same as abs).
torch.acos(input, out) Element-wise arccosine.
torch.arccos(input, out) Element-wise arccosine (same as acos).
torch.acosh(input, out) Element-wise inverse hyperbolic cosine.
torch.arccosh(input, out) Element-wise inverse hyperbolic cosine (same as acosh).
torch.add(input, other, alpha, out) Element-wise addition (with optional alpha scaling).
torch.addcdiv(input, tensor1, tensor2, value, out) Executes input + value * (tensor1 / tensor2).
torch.addcmul(input, tensor1, tensor2, value, out) Executes input + value * (tensor1 * tensor2).
torch.angle(input, out) Returns the phase angle of a complex tensor.
torch.asin(input, out) Element-wise arcsine.
torch.arcsin(input, out) Element-wise arcsine (same as asin).
torch.asinh(input, out) Element-wise inverse hyperbolic sine.
torch.arcsinh(input, out) Element-wise inverse hyperbolic sine (same as asinh).
torch.atan(input, out) Element-wise arctangent.
torch.arctan(input, out) Element-wise arctangent (same as atan).
torch.atan2(input, other, out) Element-wise two-argument arctangent.
torch.arctan2(input, other, out) Element-wise two-argument arctangent (same as atan2).
torch.atanh(input, out) Element-wise inverse hyperbolic tangent.
torch.arctanh(input, out) Element-wise inverse hyperbolic tangent (same as atanh).
torch.bitwise_not(input, out) Element-wise bitwise NOT.
torch.bitwise_and(input, other, out) Element-wise bitwise AND.
torch.bitwise_or(input, other, out) Element-wise bitwise OR.
torch.bitwise_xor(input, other, out) Element-wise bitwise XOR.
torch.bitwise_left_shift(input, other, out) Element-wise left shift.
torch.bitwise_right_shift(input, other, out) Element-wise right shift.
torch.ceil(input, out) Element-wise ceil.
torch.clamp(input, min, max, out) Clamps tensor values to a specified range.
torch.clip(input, min, max, out) Clamps tensor values to a specified range (same as clamp).
torch.conj_physical(input, out) Element-wise physical conjugate.
torch.copysign(input, other, out) Element-wise copy sign.
torch.cos(input, out) Element-wise cosine.
torch.cosh(input, out) Element-wise hyperbolic cosine.
torch.deg2rad(input, out) Converts angles from degrees to radians.
torch.div(input, other, rounding_mode, out) Element-wise division.
torch.divide(input, other, rounding_mode, out) Element-wise division (same as div).
torch.digamma(input, out) Element-wise digamma function (logarithmic derivative).
torch.erf(input, out) Element-wise error function.
torch.erfc(input, out) Element-wise complementary error function.
torch.erfinv(input, out) Element-wise inverse error function.
torch.exp(input, out) Element-wise exponential function.
torch.exp2(input, out) Element-wise power of 2.
torch.expm1(input, out) Element-wise exp(x) - 1.
torch.fake_quantize_per_channel_affine(input, scale, zero_point, axis, quant_min, quant_max) Simulates per-channel quantization.
torch.fake_quantize_per_tensor_affine(input, scale, zero_point, quant_min, quant_max) Simulates per-tensor quantization.
torch.fix(input, out) Element-wise integer part (truncation toward zero).
torch.float_power(input, exponent, out) Element-wise floating-point power.
torch.floor(input, out) Element-wise floor.
torch.floor_divide(input, other, out) Element-wise integer division.
torch.fmod(input, other, out) Element-wise modulo (remainder).
torch.frac(input, out) Element-wise fractional part.
torch.frexp(input, out) Decomposes floating-point numbers into mantissa and exponent.
torch.gradient(input, dim, spacing, edge_order) Computes the gradient of a tensor.
torch.imag(input, out) Returns the imaginary part of a complex tensor.
torch.ldexp(input, other, out) Element-wise computes input * 2**other.
torch.lerp(input, end, weight, out) Element-wise linear interpolation.
torch.lgamma(input, out) Element-wise logarithm of the gamma function.
torch.log(input, out) Element-wise natural logarithm.
torch.log10(input, out) Element-wise base-10 logarithm.
torch.log1p(input, out) Element-wise log(1 + x).
torch.log2(input, out) Element-wise base-2 logarithm.
torch.logaddexp(input, other, out) Element-wise log(exp(input) + exp(other)).
torch.logaddexp2(input, other, out) Element-wise log2(2**input + 2**other).
torch.logical_and(input, other, out) Element-wise logical AND.
torch.logical_not(input, out) Element-wise logical NOT.
torch.logical_or(input, other, out) Element-wise logical OR.
torch.logical_xor(input, other, out) Element-wise logical XOR.
torch.logit(input, eps, out) Element-wise logit function.
torch.hypot(input, other, out) Element-wise hypot function sqrt(input^2 + other^2).
torch.i0(input, out) Element-wise modified Bessel function (first kind, order 0).
torch.igamma(input, other, out) Element-wise incomplete gamma function.
torch.igammac(input, other, out) Element-wise complementary incomplete gamma function.
torch.mul(input, other, out) Element-wise multiplication.
torch.multiply(input, other, out) Element-wise multiplication (same as mul).
torch.mvlgamma(input, p, out) Element-wise logarithm of the multivariate gamma function.
torch.nan_to_num(input, nan, posinf, neginf, out) Replaces NaN with a specified value.
torch.neg(input, out) Element-wise negation.
torch.negative(input, out) Element-wise negation (same as neg).
torch.nextafter(input, other, out) Element-wise returns the next representable floating-point number.
torch.polygamma(input, n, out) Element-wise polygamma function.
torch.positive(input, out) Element-wise positive.
torch.pow(input, exponent, out) Element-wise power.
torch.quantized_batch_norm(input, weight, bias, mean, var, eps, output_scale, output_zero_point) Quantized batch normalization.
torch.quantized_max_pool1d(input, kernel_size, stride, padding, dilation, ceil_mode) Quantized max pooling (1D).
torch.quantized_max_pool2d(input, kernel_size, stride, padding, dilation, ceil_mode) Quantized max pooling (2D).
torch.rad2deg(input, out) Converts angles from radians to degrees.
torch.real(input, out) Returns the real part of a complex tensor.
torch.reciprocal(input, out) Element-wise reciprocal.
torch.remainder(input, other, out) Element-wise remainder.
torch.round(input, decimals, out) Element-wise round to nearest integer.
torch.rsqrt(input, out) Element-wise reciprocal square root.
torch.sigmoid(input, out) Element-wise sigmoid function.
torch.sign(input, out) Element-wise returns the sign (-1, 0, 1).
torch.sgn(input, out) Element-wise returns the sign vector.
torch.signbit(input, out) Element-wise checks the sign bit.
torch.sin(input, out) Element-wise sine.
torch.sinc(input, out) Element-wise sinc function sin(pi*x)/(pi*x).
torch.sinh(input, out) Element-wise hyperbolic sine.
torch.softmax(input, dim, dtype) Element-wise softmax function.
torch.sqrt(input, out) Element-wise square root.
torch.square(input, out) Element-wise square.
torch.sub(input, other, alpha, out) Element-wise subtraction.
torch.subtract(input, other, alpha, out) Element-wise subtraction (same as sub).
torch.tan(input, out) Element-wise tangent.
torch.tanh(input, out) Element-wise hyperbolic tangent.
torch.true_divide(input, other, out) Element-wise true division.
torch.trunc(input, out) Element-wise truncation (integer part).
torch.xlogy(input, other, out) Element-wise computes input * log(other).

Mathematical operations - reduction operations

Function Description
torch.argmax(input, dim, keepdim) Returns the index of the maximum value along a dimension.
torch.argmin(input, dim, keepdim) Returns the index of the minimum value along a dimension.
torch.amax(input, dim, keepdim, out) Returns the maximum value along a dimension.
torch.amin(input, dim, keepdim, out) Returns the minimum value along a dimension.
torch.aminmax(input, dim, keepdim, out) Returns the minimum and maximum values along a dimension.
torch.all(input, dim, keepdim, out) Tests whether all elements are True.
torch.any(input, dim, keepdim, out) Tests whether any elements are True.
torch.max(input, dim, keepdim, out) Computes the maximum along a specified dimension.
torch.min(input, dim, keepdim, out) Computes the minimum along a specified dimension.
torch.dist(input, other, p) Computes the p-norm distance between two tensors.
torch.logsumexp(input, dim, keepdim, out) Computes log-sum-exp.
torch.mean(input, dim, keepdim, out) Computes the mean along a specified dimension.
torch.nanmean(input, dim, keepdim, out) Computes the mean along a specified dimension (ignoring NaNs).
torch.median(input, dim, keepdim, out) Computes the median along a specified dimension.
torch.nanmedian(input, dim, keepdim, out) Computes the median along a specified dimension (ignoring NaNs).
torch.mode(input, dim, keepdim, out) Computes the mode along a specified dimension.
torch.norm(input, p, dim, keepdim, out) Computes the p-norm.
torch.nansum(input, dim, keepdim, out) Computes the sum along a specified dimension (ignoring NaNs).
torch.prod(input, dim, keepdim, dtype, out) Computes the product along a specified dimension.
torch.quantile(input, q, dim, keepdim, out, method) Computes the quantile.
torch.nanquantile(input, q, dim, keepdim, out, method) Computes the quantile (ignoring NaNs).
torch.std(input, dim, unbiased, keepdim, out) Computes the standard deviation.
torch.std_mean(input, dim, unbiased, keepdim) Computes the standard deviation and mean.
torch.sum(input, dim, keepdim, dtype, out) Computes the sum along a specified dimension.
torch.unique(input, sorted, return_inverse, return_counts, dim) Returns unique elements.
torch.unique_consecutive(input, sorted, return_inverse, return_counts, dim) Returns consecutive unique elements.
torch.var(input, dim, unbiased, keepdim, out) Computes the variance.
torch.var_mean(input, dim, unbiased, keepdim) Computes the variance and mean.
torch.count_nonzero(input, dim) Counts the number of nonzero elements.
torch.hash_tensor(input) Computes the hash value of a tensor.

Mathematical operations - comparison operations

Function Description
torch.allclose(input, other, rtol, atol, equal_nan) Checks whether two tensors are close (all elements).
torch.argsort(input, dim, descending, stable) Returns the indices that sort the tensor.
torch.eq(input, other, out) Element-wise equality comparison.
torch.equal(input, other) Checks whether two tensors are exactly equal.
torch.ge(input, other, out) Element-wise greater-than-or-equal comparison.
torch.greater_equal(input, other, out) Element-wise greater-than-or-equal comparison (same as ge).
torch.gt(input, other, out) Element-wise greater-than comparison.
torch.greater(input, other, out) Element-wise greater-than comparison (same as gt).
torch.isclose(input, other, rtol, atol, equal_nan) Checks whether two tensors are close (element-wise).
torch.isfinite(input, out) Checks whether values are finite.
torch.isin(elements, test_elements, assume_unique, invert) Checks whether elements are in a set.
torch.isinf(input, out) Checks whether values are infinite.
torch.isposinf(input, out) Checks whether values are positive infinity.
torch.isneginf(input, out) Checks whether values are negative infinity.
torch.isnan(input, out) Checks whether values are NaN.
torch.isreal(input, out) Checks whether values are real.
torch.kthvalue(input, k, dim, keepdim, out) Returns the k-th smallest element and its index.
torch.le(input, other, out) Element-wise less-than-or-equal comparison.
torch.less_equal(input, other, out) Element-wise less-than-or-equal comparison (same as le).
torch.lt(input, other, out) Element-wise less-than comparison.
torch.less(input, other, out) Element-wise less-than comparison (same as lt).
torch.maximum(input, other, out) Element-wise maximum.
torch.minimum(input, other, out) Element-wise minimum.
torch.fmax(input, other, out) Element-wise maximum (ignoring NaNs).
torch.fmin(input, other, out) Element-wise minimum (ignoring NaNs).
torch.ne(input, other, out) Element-wise inequality comparison.
torch.not_equal(input, other, out) Element-wise inequality comparison (same as ne).
torch.sort(input, dim, descending, stable, out) Sorts along a specified dimension.
torch.topk(input, k, dim, largest, sorted, out) Returns the largest k elements and their indices.
torch.msort(input, out) Sorts along the last dimension (returns the sorted tensor).

Mathematical operations - spectral operations

Function Description
torch.stft(input, n_fft, hop_length, win_length, window, center, normalized, onesided, return_complex) Short-time Fourier transform.
torch.istft(input, n_fft, hop_length, win_length, window, center, normalized, onesided, length, return_complex) Inverse short-time Fourier transform.
torch.bartlett_window(window_length, periodic, dtype, device) Bartlett window.
torch.blackman_window(window_length, periodic, dtype, device) Blackman window.
torch.hamming_window(window_length, periodic, alpha, beta, dtype, device) Hamming window.
torch.hann_window(window_length, periodic, dtype, device) Hann window.
torch.kaiser_window(window_length, periodic, beta, dtype, device) Kaiser window.

Mathematical Operations - Other Operations

Function Description
torch.atleast_1d(*tensors) Converts the input to a tensor of at least 1 dimension.
torch.atleast_2d(*tensors) Converts the input to a tensor of at least 2 dimensions.
torch.atleast_3d(*tensors) Converts the input to a tensor of at least 3 dimensions.
torch.bincount(input, weights, minlength) Counts the occurrences of nonnegative integers.
torch.block_diag(*tensors) Constructs a block diagonal matrix from the input tensor.
torch.broadcast_tensors(*tensors) Broadcasts the input to a common shape.
torch.broadcast_to(input, shape) Broadcasts the tensor to a specified shape.
torch.broadcast_shapes(*shapes) Broadcasts shapes to be compatible for operations.
torch.bucketize(input, boundaries, right) Maps the input to bucket indices.
torch.cartesian_prod(*tensors) Computes the Cartesian product.
torch.cdist(x1, x2, p, compute_mode) Computes pairwise distances.
torch.clone(input, memory_format) Returns a copy of the tensor.
torch.combinations(input, r, with_replacement) Computes combinations.
torch.corrcoef(input) Computes the correlation coefficient matrix.
torch.cov(input, correction, fweights, aweights) Computes the covariance matrix.
torch.cross(input, dim, out) Computes the cross product.
torch.cummax(input, dim, out) Accumulates the maximum value along a dimension.
torch.cummin(input, dim, out) Accumulates the minimum value along a dimension.
torch.cumprod(input, dim, out) Accumulates the product along a dimension.
torch.cumsum(input, dim, out, dtype) Accumulates the sum along a dimension.
torch.diag(input, diagonal, out) Creates a diagonal matrix or extracts the diagonal.
torch.diag_embed(input, offset, dim1, dim2, out) Embeds the input as a diagonal.
torch.diagflat(input, offset, out) Creates a diagonal matrix (flattened input).
torch.diagonal(input, offset, dim1, dim2, out) Extracts diagonal elements.
torch.diff(input, n, dim, prepend, append, out) Computes the difference.
torch.einsum(equation, *operands) Einstein summation convention.
torch.flatten(input, start_dim, end_dim, out) Flattens the tensor.
torch.ravel(input, out) Flattens to a one-dimensional tensor.
torch.kron(input, other, out) Computes the Kronecker product.
torch.meshgrid(*tensors, indexing) Creates a grid.
torch.lcm(input, other, out) Element-wise least common multiple.
torch.logcumsumexp(input, dim, out) Cumulative log-sum-exp along a dimension.
torch.renorm(input, p, dim, maxnorm, out) Renormalizes to a specified norm.
torch.roll(input, shifts, dims) Rolls tensor elements.
torch.searchsorted(sorted_sequence, values, side, sorter) Searches for positions in a sorted sequence.
torch.tensordot(a, b, dims) Computes the tensor dot product.
torch.trace(input, out) Computes the trace of a matrix.
torch.tril(input, diagonal, out) Extracts the lower triangular matrix.
torch.tril_indices(row, column, offset, dtype, device, layout) Generates lower triangular indices.
torch.triu(input, diagonal, out) Extracts the upper triangular matrix.
torch.triu_indices(row, column, offset, dtype, device, layout) Generates upper triangular indices.
torch.unflatten(input, dim, sizes) Unfolds the tensor.
torch.vander(x, N, increasing, out) Creates a Vandermonde matrix.

Linear Algebra (BLAS and LAPACK)

Function Description
torch.addbmm(input, batch1, batch2, beta, alpha, out) Batch matrix-matrix multiply-add.
torch.addmm(input, mat1, mat2, beta, alpha, out) Matrix multiply-add.
torch.addmv(input, mat, vec, beta, alpha, out) Matrix-vector multiply-add.
torch.addr(input, vec1, vec2, beta, alpha, out) Vector outer product add.
torch.baddbmm(input, batch1, batch2, beta, alpha, out) Batch matrix multiply-add (bmm + add).
torch.bmm(input, mat2, out) Batch matrix multiplication.
torch.chain_matmul(*matrices) Chained matrix multiplication.
torch.cholesky(input, upper, out) Cholesky decomposition.
torch.cholesky_inverse(input, upper, out) Cholesky decomposition inverse.
torch.cholesky_solve(input, input2, upper, out) Cholesky decomposition solving linear equations.
torch.dot(input, other, out) Computes the dot product of two vectors.
torch.geqrf(input, out) QR decomposition (geqrf).
torch.ger(input, other, out) Computes the outer product of vectors.
torch.inner(input, other, out) Computes the inner product.
torch.inverse(input, out) Computes the inverse of a matrix.
torch.det(input, out) Computes the determinant of a matrix.
torch.logdet(input, out) Computes the logarithm of the determinant.
torch.slogdet(input, out) Computes the sign and log absolute value of the determinant.
torch.lu(input, pivot, get_infos, out) LU decomposition.
torch.lu_solve(input, LU_data, LU_pivots, out) LU decomposition solve.
torch.lu_unpack(LU_data, LU_pivots, unpack_data, unpack_pivots) Unpacks the LU decomposition result.
torch.matmul(input, other, out) Matrix multiplication (supports different dimensions).
torch.matrix_power(input, n, out) Matrix power.
torch.matrix_exp(input, out) Matrix exponential.
torch.mm(input, mat2, out) Matrix multiplication (2D).
torch.mv(input, vec, out) Matrix-vector multiplication.
torch.orgqr(input, q, out) Reconstructs Q from QR decomposition.
torch.ormqr(input, mat, vec, left, transpose, out) ormqr operation.
torch.outer(input, other, out) Computes the outer product of vectors.
torch.pinverse(input, rcond, out) Computes the Moore-Penrose pseudoinverse.
torch.qr(input, out) QR decomposition.
torch.svd(input, some, compute_uv, out) Singular value decomposition.
torch.svd_lowrank(input, q, niter, M) Low-rank SVD approximation.
torch.pca_lowrank(input, q, center, niter) Low-rank PCA approximation.
torch.lobpcg(input, K, B, X, M, P, max_iter, tol, debug, ortho_iparams, fpfloor) LOBPCG eigenvalue solver.
torch.trapz(y, x, dim, out) Trapezoidal integration (deprecated, use trapezoid).
torch.trapezoid(y, x, dim, out) Trapezoidal integration.
torch.cumulative_trapezoid(y, x, dim, out) Cumulative trapezoidal integration.
torch.triangular_solve(input, A, upper, transpose, unitriangular, out) Triangular matrix solve.
torch.vdot(input, other, out) Computes the dot product of vectors (complex-aware).

Device Management

Function Description
torch.cuda.is_available() Checks whether CUDA is available.
torch.cuda.device_count() Returns the number of CUDA devices.
torch.cuda.current_device() Returns the current CUDA device index.
torch.cuda.device(name) Creates a device object.
torch.cuda.device_context(device) Creates a device context.
torch.device(device) Creates a device object (e.g.,'cpu'or'cuda:0')。
torch.Tensor.to(device) Moves the tensor to the specified device.
torch.get_device_module(device_type) Gets the device module (e.g., cuda, mps).

Parallel Computing

Function Description
torch.get_num_threads() Gets the total number of threads for CPU operations.
torch.set_num_threads(int) Sets the number of threads for CPU operations.
torch.get_num_interop_threads() Gets the number of inter-op parallel threads.
torch.set_num_interop_threads(int) Sets the number of inter-op parallel threads.

Utility Functions

Function Description
torch.compiled_with_cxx11_abi() Checks whether compiled with C++11 ABI.
torch.result_type(tensor, other) Returns the dtype of the operation result.
torch.can_cast(from_dtype, to_dtype) Checks whether data types can be converted.
torch.promote_types(type1, type2) Returns the promoted data type.
torch.use_deterministic_algorithms(mode, warn_only) Enables/disables deterministic algorithms.
torch.are_deterministic_algorithms_enabled() Checks whether deterministic algorithms are enabled.
torch.is_deterministic_algorithms_warn_only_enabled() Checks whether deterministic algorithms are in warning mode.
torch.set_deterministic_debug_mode(debug_mode) Sets deterministic debug mode.
torch.get_deterministic_debug_mode() Gets deterministic debug mode.
torch.set_float32_matmul_precision(precision) Sets the precision for float32 matrix multiplication.
torch.get_float32_matmul_precision() Gets the precision for float32 matrix multiplication.
torch.set_warn_always(enabled) Sets whether to always display warnings.
torch.is_warn_always_enabled() Checks whether warnings are always displayed.
torch.vmap(fn, in_dims, out_dims, randomness, chunk_size) Vectorized mapping.
torch._assert(condition, message) Assertion check (internal use).
torch.typename(t) Returns the string representation of the type.

Compile Optimization

Function Description
torch.compile(model, backend, options, dynamic) Compiles a PyTorch model for optimization.

Example

Example

import torch

# Create a tensor
x = torch.tensor([1, 2, 3])
y = torch.zeros(2, 3)

# Mathematical operations
z = torch.add(x, 1)  # Element-wise add 1
print(z)

# Indexing and slicing
mask = x > 1
selected = torch.masked_select(x, mask)
print(selected)

# Device management
if torch.cuda.is_available():
    device = torch.device('cuda')
    x = x.to(device)
    print(x.device)

# Matrix operations
a = torch.randn(3, 4)
b = torch.randn(4, 5)
c = torch.matmul(a, b)
print(c.shape)

# Gradient computation
x = torch.tensor([1., 2., 3.], requires_grad=True)
y = x.sum()
y.backward()
print(x.grad)

Output:

tensor([2, 3, 4])
tensor([2, 3])

For more detailed information, please refer toPyTorch Official Documentation。

Other Extensions