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

如何在PyTorch中高效将RGB图像编码为n类One-Hot张量?

Efficient One-Hot Encoding for RGB Segmentation Labels (No Nested Loops!)

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.

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:

  1. Reshape the input image to a flat list of RGB pixels
  2. Use broadcasting to compare each pixel to all colors in your mapping
  3. Find the matching class for each pixel
  4. 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 True per row in matches).
  • 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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.13 09:13:37