如何在PyTorch中创建切片掩码?能否基于切片操作指定掩码?
Absolutely! Generating a mask that aligns with a specific slice of your tensor is totally doable in PyTorch—let's break it down using your exact example.
First, let's set up the tensor you mentioned:
import torch A = torch.arange(6).view((2,3)) # Result: tensor([[0, 1, 2], [3, 4, 5]])
Method 1: Initialize Zero Mask + Slice Assignment
The simplest approach is to start with a mask filled entirely with 0s (matching the shape of A), then set the positions corresponding to your target slice to 1s:
# Create a mask of 0s with the same shape and integer dtype as A mask_slice = torch.zeros_like(A, dtype=torch.int) # Assign 1s to the slice A[:, 1:] mask_slice[:, 1:] = 1 print(mask_slice) # Output: tensor([[0, 1, 1], [0, 1, 1]])
Method 2: Boolean Mask Conversion
If you prefer working with boolean tensors (which are common in PyTorch for indexing), you can create a boolean mask first and then convert it to integers (where True becomes 1 and False becomes 0):
# Initialize a boolean mask filled with False bool_mask = torch.zeros_like(A, dtype=torch.bool) # Mark the target slice as True bool_mask[:, 1:] = True # Convert boolean values to integers mask_slice = bool_mask.int() print(mask_slice) # Same output as before: tensor([[0, 1, 1], [0, 1, 1]])
Handling More Complex Slices
This approach works for any slice you can define in PyTorch. For example, if you wanted a mask for A[1:, 0:2] (the last row, first two columns), you'd do:
mask_complex = torch.zeros_like(A, dtype=torch.int) mask_complex[1:, 0:2] = 1 print(mask_complex) # Output: tensor([[0, 0, 0], [1, 1, 0]])
The core idea is straightforward: start with a base mask (all 0s or all 1s, depending on your needs), then use standard PyTorch slicing syntax to update the regions you want to highlight in the mask.
内容的提问来源于stack exchange,提问作者DsCpp

