You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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 from nc to the last column (exclusive of n, which matches your slice self.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 in self.X).
  • scatter(1, sliced_index, src): Along dimension 1 (columns), this places each element from src into the position specified by sliced_index in the original self.X tensor. 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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.07 08:12:43