torch.unique()的dim参数作用及矩阵行内去重置零实现咨询
torch.unique(dim) Let's tackle your two questions one by one—first solving the deduplication task without for loops, then breaking down how torch.unique() uses the dim parameter.
1. Row-wise Deduplication (Set Duplicates to 0)
To keep only the first occurrence of each value in every row and set subsequent duplicates to 0, we can use PyTorch's broadcasting and boolean masking to vectorize the operation (no loops needed):
import torch # Your input tensor input_tensor = torch.Tensor([[1, 2, 3, 4, 3, 3, 4], [1, 6, 3, 5, 3, 5, 4]]) # Step 1: Expand dimensions to enable row-wise element comparison expanded = input_tensor.unsqueeze(2) # Shape: (2, 7, 1) # Compare each element to all elements in its row (broadcasting) equal_mask = expanded == input_tensor.unsqueeze(1) # Shape: (2, 7, 7) # Step 2: Count cumulative occurrences of each element up to its position counts = torch.cumsum(equal_mask, dim=1) # Shape: (2, 7, 7) # For each element, check if it's the first occurrence (count == 1) first_occurrence_mask = counts.diagonal(dim1=1, dim2=2) == 1 # Shape: (2, 7) # Step 3: Keep first occurrences, set duplicates to 0 result = torch.where(first_occurrence_mask, input_tensor, torch.zeros_like(input_tensor)) print(result)
Output:
tensor([[1., 2., 3., 4., 0., 0., 0.], [1., 6., 3., 5., 0., 0., 4.]])
How this works:
- We use broadcasting to compare every element in a row to every other element in the same row, creating a boolean matrix where
equal_mask[i,j,k]isTrueifinput_tensor[i,j] == input_tensor[i,k]. torch.cumsumalongdim=1counts how many times the element at positionjhas appeared up to that point in the row.- We extract the diagonal of the counts matrix (since we only care about the count for each element's own position) and mask for values equal to 1 (first occurrence).
- Finally,
torch.wherereplaces non-first-occurrence elements with 0.
2. What Does torch.unique(dim) Actually Do?
Your confusion makes total sense—torch.unique(dim) doesn't do element-wise deduplication within rows/columns. Instead, it treats entire slices along the specified dimension as single "elements" and finds unique slices.
For example, when you run:
output = torch.unique(torch.Tensor([[4,2,52,2,2],[5,2,6,6,5]]), dim=1)
dim=1tells PyTorch to look at columns (each column is a 2-element slice) as the units to deduplicate.- The input has 5 columns:
[4,5],[2,2],[52,6],[2,6],[2,5]—all of these are unique, so none are removed. - By default,
torch.uniquesorts the unique slices (sorted=True), so it orders the columns by their first element, then second. That's why you get the output:
(These are the sorted unique columns:tensor([[ 2., 2., 2., 4., 52.], [ 2., 5., 6., 5., 6.]])[2,2],[2,5],[2,6],[4,5],[52,6].)
If you used dim=0, PyTorch would treat rows as the units to deduplicate, finding unique rows in your tensor.
内容的提问来源于stack exchange,提问作者Sean Lee

