PyTorch中基于索引-值向量的重复元素向量向量化加法优化
Great question! Using index-value pairs to compress high-duplication vectors is a smart move for memory efficiency, and optimizing their addition in PyTorch is totally worth it—let’s break down two solid, efficient approaches that leverage PyTorch’s native vectorized operations (no slow loops required).
Approach 1: Vectorized Segmentation (Memory-Efficient for Large Vectors)
This method works entirely in the compressed space, avoiding expanding the full vector. It focuses on merging segment boundaries from both input pairs, computing sums per segment, then collapsing consecutive duplicates.
Implementation Code
import torch def add_compressed_vectors(indices1, values1, indices2, values2, total_length): # Merge, deduplicate, and sort all boundary indices (plus total length as end marker) combined_indices = torch.cat([indices1, indices2, torch.tensor([total_length], device=indices1.device)]).unique(sorted=True) # Split into segment start/end points segment_starts = combined_indices[:-1] segment_ends = combined_indices[1:] # Find the corresponding value from each input for every segment start idx1 = torch.searchsorted(indices1, segment_starts, right=True) - 1 val1 = values1[idx1] idx2 = torch.searchsorted(indices2, segment_starts, right=True) - 1 val2 = values2[idx2] # Calculate sum for each segment segment_values = val1 + val2 # Collapse consecutive segments with identical sums change_mask = torch.cat([torch.tensor([True], device=segment_values.device), segment_values[1:] != segment_values[:-1]]) final_indices = segment_starts[change_mask] final_values = segment_values[change_mask] return final_indices, final_values
Why This Works
torch.searchsorteduses efficient binary search to map segment starts to the correct values in O(log n) time per segment, way faster than looping.- All operations are vectorized, so they run efficiently on both CPU and GPU.
- We never expand the full vector, making this ideal for extremely large total lengths where memory is a constraint.
Approach 2: Sparse Tensor Conversion (Clean, Maintainable Code)
PyTorch’s sparse tensor utilities handle a lot of the heavy lifting for you. Convert your compressed pairs to COO-format sparse tensors, add them natively, then convert back to the index-value format.
Implementation Code
def compressed_to_sparse(indices, values, total_length): # Calculate length of each segment segment_lengths = torch.diff(torch.cat([indices, torch.tensor([total_length], device=indices.device)])) # Generate full indices for each segment coo_indices = torch.repeat_interleave(indices, segment_lengths) # Create COO sparse tensor return torch.sparse_coo_tensor(coo_indices.unsqueeze(0), values.repeat_interleave(segment_lengths), size=(total_length,)) def sparse_to_compressed(sparse_tensor): # Extract COO indices and values coo_indices = sparse_tensor.indices().squeeze() values = sparse_tensor.values() # Find where values change to collapse segments change_mask = torch.cat([torch.tensor([True], device=values.device), values[1:] != values[:-1]]) compressed_indices = coo_indices[change_mask] compressed_values = values[change_mask] return compressed_indices, compressed_values # Example Usage indices_a = torch.tensor([0, 3, 5]) values_a = torch.tensor([1, 7, 4]) indices_b = torch.tensor([0, 2, 6]) values_b = torch.tensor([2, 3, 1]) total_len = 7 sparse_a = compressed_to_sparse(indices_a, values_a, total_len) sparse_b = compressed_to_sparse(indices_b, values_b, total_len) sparse_sum = sparse_a + sparse_b final_indices, final_values = sparse_to_compressed(sparse_sum)
Pros & Cons
- Pros: Code is cleaner and easier to maintain—you let PyTorch handle segment merging and value alignment.
- Cons: Generating full COO indices can use more memory if your total vector length is massive (e.g., millions of elements). Stick with Approach 1 if memory is tight.
Quick Notes
- Both approaches assume your input index vectors are strictly increasing and start with
0. If your inputs don’t meet this, add a quick preprocessing step to sort, deduplicate, and prepend0if needed. - Always run operations on the same device (CPU/GPU) to avoid unnecessary data transfers—pass the device explicitly if working with multiple devices.
内容的提问来源于stack exchange,提问作者kuppern87

