torch_blue.vi.VIReturn

class torch_blue.vi.VIReturn(data: torch.Tensor, log_probs: torch.Tensor | None)

Bases: torch.Tensor

A subclass of torch.Tensor that also stores log probabilities.

A VIReturn object behaves like a torch.Tensor for all practical purposes, but provides the optional attribute log_probs. However, it should not be used as replacement since certain pytorch operation will lose the log prob information. It is almost exclusively used as the return formate for Tensor, but still provide the log prob information when needed. torch_blue losses may require this format as input.

new_empty(size: Tuple[int, ...], **kwargs: Any) VIReturn

Return a VIReturn of size size filled with uninitialized data.

clone(*args: Any, **kwargs: Any) VIReturn

clone(*, memory_format=torch.preserve_format) -> Tensor

See torch.clone()

to(*args: Any, **kwargs: Any) VIReturn

to(*args, **kwargs) -> Tensor

Performs Tensor dtype and/or device conversion. A torch.dtype and torch.device are inferred from the arguments of self.to(*args, **kwargs).

Note

If the self Tensor already has the correct torch.dtype and torch.device, then self is returned. Otherwise, the returned tensor is a copy of self with the desired torch.dtype and torch.device.

Note

If self requires gradients (requires_grad=True) but the target dtype specified is an integer type, the returned tensor will implicitly set requires_grad=False. This is because only tensors with floating-point or complex dtypes can require gradients.

Here are the ways to call to:

to(dtype, non_blocking=False, copy=False, memory_format=torch.preserve_format) Tensor

Returns a Tensor with the specified dtype

Args:

memory_format (torch.memory_format, optional): the desired memory format of returned Tensor. Default: torch.preserve_format.

Note

According to C++ type conversion rules, converting floating point value to integer type will truncate the fractional part. If the truncated value cannot fit into the target type (e.g., casting torch.inf to torch.long), the behavior is undefined and the result may vary across platforms.

to(device=None, dtype=None, non_blocking=False, copy=False, memory_format=torch.preserve_format) Tensor

Returns a Tensor with the specified device and (optional) dtype. If dtype is None it is inferred to be self.dtype. When non_blocking is set to True, the function attempts to perform the conversion asynchronously with respect to the host, if possible. This asynchronous behavior applies to both pinned and pageable memory. However, caution is advised when using this feature. For more information, refer to the tutorial on good usage of non_blocking and pin_memory. When copy is set, a new Tensor is created even when the Tensor already matches the desired conversion.

Args:

memory_format (torch.memory_format, optional): the desired memory format of returned Tensor. Default: torch.preserve_format.

to(other, non_blocking=False, copy=False) Tensor

Returns a Tensor with same torch.dtype and torch.device as the Tensor other. When non_blocking is set to True, the function attempts to perform the conversion asynchronously with respect to the host, if possible. This asynchronous behavior applies to both pinned and pageable memory. However, caution is advised when using this feature. For more information, refer to the tutorial on good usage of non_blocking and pin_memory. When copy is set, a new Tensor is created even when the Tensor already matches the desired conversion.

Example:

>>> tensor = torch.randn(2, 2)  # Initially dtype=float32, device=cpu
>>> tensor.to(torch.float64)
tensor([[-0.5044,  0.0005],
        [ 0.3310, -0.0584]], dtype=torch.float64)

>>> cuda0 = torch.device('cuda:0')
>>> tensor.to(cuda0)
tensor([[-0.5044,  0.0005],
        [ 0.3310, -0.0584]], device='cuda:0')

>>> tensor.to(cuda0, dtype=torch.float64)
tensor([[-0.5044,  0.0005],
        [ 0.3310, -0.0584]], dtype=torch.float64, device='cuda:0')

>>> other = torch.randn((), dtype=torch.float64, device=cuda0)
>>> tensor.to(other, non_blocking=True)
tensor([[-0.5044,  0.0005],
        [ 0.3310, -0.0584]], dtype=torch.float64, device='cuda:0')
static from_tensor(input_: torch.Tensor, log_probs: torch.Tensor | None) VIReturn

Turn a torch.Tensor into a VIReturn.

This is an inplace operation and does not copy data.

Parameters:
  • input_ (Tensor) – The Tensor to convert.

  • log_probs (Optional[Tensor]) – The log probabilities to attach to input_.

Returns:

The input converted to a VIReturn with the specified log probabilities.

Return type:

VIReturn