如何在PyTorch中高效将RGB图像编码为n类One-Hot张量?
Great question! Ditching those slow nested loops is a crucial optimization when preparing labels for U-Net—vectorized operations in NumPy will save you tons of time, especially with large datasets. Let's walk through how to rebuild your HotEncoder class for speed, using your existing color-to-class mapping.
First: Fix Your Color Mapping (Optional but Recommended)
Quick note: One-Hot matrices typically use 0-indexed classes (since np.eye defaults to this). Your binary mode currently maps colors to 1 and 2—you might want to adjust this to 0 and 1 to avoid extra unused channels. But we’ll work with either setup in the code below.
Step 1: Optimize Color Palette Generation
Your current gen_colors method uses nested loops to count unique colors—let’s replace that with np.unique, which is way faster for this task:
import glob from tqdm import tqdm import numpy as np import skimage.io class HotEncoder(): def __init__(self, dir, extension, is_binary=True): self.dir = dir self.extension = extension self.is_binary = is_binary if is_binary: # Adjust to 0-indexed if preferred: {(0, 0, 0): 0, (255, 255, 255): 1} self.color = {(0, 0, 0): 1, (255, 255, 255): 2} else: self.color = dict() def gen_colors(self): """Iterates through the dataset to find unique colors for one-hot encoding""" if self.is_binary: return self.color else: images = glob.glob(self.dir + '/*.' + self.extension) all_pixels = [] for img in tqdm(images, desc="Generating Color Palette"): image = skimage.io.imread(img) # Reshape image to (num_pixels, 3) to collect all RGB values pixels = image.reshape(-1, 3) all_pixels.append(pixels) # Combine all pixels and find unique RGB tuples all_pixels = np.concatenate(all_pixels, axis=0) unique_colors = np.unique(all_pixels, axis=0) # Build color-to-class mapping (0-indexed by default) self.color = {tuple(clr): idx for idx, clr in enumerate(unique_colors)} return self.color
Step 2: Vectorized One-Hot Label Generation
Now the main event: replacing your PerPixelClassMatrix with a loop-free vectorized approach. Here’s how it works:
- Reshape the input image to a flat list of RGB pixels
- Use broadcasting to compare each pixel to all colors in your mapping
- Find the matching class for each pixel
- Convert the class matrix to a One-Hot matrix
Add this method to your HotEncoder class:
def get_one_hot_label(self, image): """Generates an n-channel One-Hot matrix from an RGB segmentation mask""" # Get image dimensions and reshape to (num_pixels, 3) height, width, _ = image.shape flat_pixels = image.reshape(-1, 3) # Convert color mapping to NumPy arrays for broadcasting color_list = np.array(list(self.color.keys())) class_list = np.array(list(self.color.values())) # Compare each flat pixel to all colors (broadcasting magic) # Result is a (num_pixels, num_colors) boolean matrix matches = np.all(flat_pixels[:, np.newaxis] == color_list, axis=2) # Get the class index for each pixel class_indices = class_list[np.argmax(matches, axis=1)] # Convert to One-Hot matrix: (height, width, num_classes) num_classes = len(self.color) one_hot = np.eye(num_classes)[class_indices].reshape(height, width, num_classes) return one_hot
Key Details
- Broadcasting:
flat_pixels[:, np.newaxis]adds a dimension to the pixel array, letting us compare every pixel to every color in a single operation (no loops!). - np.argmax: Finds which color matches each pixel (since valid masks will have exactly one
Trueper row inmatches). - np.eye: Converts class indices to a One-Hot matrix efficiently, leveraging NumPy’s optimized C backend.
Bonus: Handle Unseen Colors
If you might encounter colors not in your mapping, add a fallback to assign them to a default class (e.g., background):
# Add this after computing `matches` # Check if any color matched for each pixel has_match = np.any(matches, axis=1) # Assign unseen pixels to class 0 (or another value of your choice) class_indices[~has_match] = 0
Performance Boost
For a 1024x1024 image, this vectorized approach will run 100-1000x faster than nested Python loops—critical when processing large datasets for U-Net training.
内容的提问来源于stack exchange,提问作者user8097625

