PyTorch torch.result_type Function


Pytorch torch 参考手册Pytorch torch Reference Manual

torch.result_typeis a function in PyTorch used to determine the result type of an operation. It accepts tensors or scalars as input and returns the data type that should be used after performing the relevant operation.

Function Definition

torch.result_type(tensor, scalar)

Parameter Description

  • tensor: input tensor or data type
  • scalar: scalar value or another tensor

Usage Example

Example

import torch

# Create tensors of different types
a = torch.tensor([1, 2, 3], dtype=torch.float32)
b = torch.tensor([4, 5, 6], dtype=torch.float64)

# Get the result type
result_dtype = torch.result_type(a, b)

print("Result type of float32 and float64:", result_dtype)

# Using a scalar
c = torch.tensor([1, 2, 3], dtype=torch.int32)
result_dtype2 = torch.result_type(c, 1.5)

print("Result type of int32 and 1.5:", result_dtype2)

The output result is:

float32 和 float64 的结果类型: torch.float64
int32 和 1.5 的结果类型: torch.float32

Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions