PyTorch中如何将张量原地赋值转为非原地scatter操作
Solution
To convert your in-place assignment to a non-in-place scatter operation, you need to create an index tensor that specifies exactly which columns in self.X should be replaced by your source tensor (let's call it src, which is D in your minimal example or mm(S_, Z[:, :n - nc]) in your original code).
Step-by-Step Implementation
First, let's walk through your minimal example to see how it works:
import torch # Your minimal example inputs a = torch.tensor([[0, 1, 2], [3, 4, 5]]) D = torch.tensor([[6], [7]]) nc = 2 # Since we're replacing columns starting at index 2 # 1. Get the total number of columns in your tensor n = a.size(1) # 2. Create the column indices we want to replace: from nc to n-1 base_indices = torch.arange(nc, n, dtype=torch.long) # tensor([2]) for n=3, nc=2 # 3. Reshape and repeat to match the batch size (number of rows) sliced_index = base_indices.unsqueeze(0).repeat(a.size(0), 1) # tensor([[2], [2]]) # 4. Perform the non-in-place scatter operation not_in_place = a.scatter(1, sliced_index, D) print(not_in_place) # Output: tensor([[0, 1, 6], [3, 4, 7]])
Generalizing to Your Original Code
For your original use case, here's how to adapt this:
# Assume self.X is your original tensor, nc is the split point n = self.X.size(1) src = torch.mm(S_, Z[:, :n - nc]) # Your source tensor # Create the index tensor sliced_index = torch.arange(nc, n, dtype=torch.long) sliced_index = sliced_index.unsqueeze(0).repeat(self.X.size(0), 1) # Ensure index is on the same device as self.X (critical for GPU usage) sliced_index = sliced_index.to(self.X.device) # Non-in-place assignment using scatter self.X = self.X.scatter(1, sliced_index, src)
Key Explanations
torch.arange(nc, n): Generates a 1D tensor containing all column indices fromncto the last column (exclusive ofn, which matches your sliceself.X[:, nc:]).unsqueeze(0).repeat(...): Turns the 1D index array into a 2D tensor where each row is identical (since we want to replace the same columns across all rows inself.X).scatter(1, sliced_index, src): Along dimension 1 (columns), this places each element fromsrcinto the position specified bysliced_indexin the originalself.Xtensor. The result is a new tensor, so this is fully non-in-place.- Device/Dtype Matching: Always make sure your index tensor uses the same device (CPU/GPU) and integer dtype (
torch.long) as your original tensor to avoid errors.
This approach perfectly replicates the behavior of your in-place assignment but without modifying the original tensor directly—instead, it creates a new tensor with the updated values.
内容的提问来源于stack exchange,提问作者DsCpp
相关产品推荐
相关产品推荐

