如何在TensorFlow中为segment_ids张量内的各分段实例分配唯一ID(不使用py_func)
Solution: Assign Per-Segment Unique IDs in TensorFlow (No
py_func) To assign a unique 0-based ID to each instance within its segment (even when segments are non-contiguous), we can use a combination of tf.unique and tf.scan to track cumulative counts per segment. Here's a step-by-step breakdown:
Approach
- Map Segment IDs to Consecutive Indices: Use
tf.uniqueto convert arbitrary segment IDs into a compact set of consecutive integers. This makes it easy to track counts for each segment using a fixed-size state tensor. - Track Cumulative Counts with
tf.scan: Process each element sequentially, maintaining a state that keeps track of how many times each segment has been seen so far. For each element, we:- Retrieve the current count for its segment.
- Append this count to our output (this is the unique ID for the instance).
- Increment the segment's count in the state.
Code Implementation
import tensorflow as tf # Example input tensor x = tf.constant([1, 1, 2, 2, 3, 3, 4, 1]) # Step 1: Get unique segments and their indices in the original tensor unique_segments, segment_indices = tf.unique(x) num_unique_segments = tf.shape(unique_segments)[0] # Step 2: Define the scan function to track counts per segment def update_count(state, segment_idx): # Get the current count for this segment current_count = state[segment_idx] # Update the state to increment the count for this segment updated_state = tf.tensor_scatter_nd_update( state, indices=[[segment_idx]], updates=[current_count + 1] ) # Return updated state and the current count (our output for this element) return updated_state, current_count # Initialize state with zeros for each unique segment initial_state = tf.zeros(num_unique_segments, dtype=tf.int32) # Run the scan operation to compute per-instance IDs _, per_segment_ids = tf.scan(update_count, segment_indices, initializer=initial_state) # Verify the result print("Input segment IDs:", x.numpy()) print("Output per-segment unique IDs:", per_segment_ids.numpy())
Output
Input segment IDs: [1 1 2 2 3 3 4 1] Output per-segment unique IDs: [0 1 0 1 0 1 0 2]
Key Notes
- No Python Overhead: This uses pure TensorFlow operations, avoiding the performance cost of
py_funcwhich switches between Python and TensorFlow execution contexts. - Handles Non-Contiguous Segments: Unlike methods that rely on contiguous blocks, this works even when segments reappear later in the tensor (like the
1segment in the example). - Works with Any Hashable Type:
tf.uniquesupports integers, strings, and other hashable tensor types, so this approach isn't limited to numeric segment IDs.
内容的提问来源于stack exchange,提问作者ma7555
相关产品推荐
相关产品推荐

