基于TensorFlow Estimator API实现图像语义分割
Hey there! Let's walk through how to refine your input_fn for TensorFlow Estimator-based semantic segmentation. Your current setup has a few key gaps that'll hold you back with larger datasets, so here's a breakdown of best practices and fixes:
Right now, you're loading all image/label paths into a NumPy array and converting it to a tf.constant—this will crash or slow to a crawl if you're working with a large dataset (like Cityscapes, which I assume you're using given the path structure). Instead, let tf.data handle the path list dynamically. If you need to keep the sorted matching between images and labels, you can still pass the sorted lists to from_tensor_slices, but avoid wrapping them in tf.constant (just pass the Python list directly—TensorFlow will handle it efficiently).
Your current code only holds file paths, not the actual image data. You need a parsing function to load and preprocess images and segmentation masks on the fly:
def parse_image_mask(image_path, mask_path): # Load and preprocess image img = tf.io.read_file(image_path) img = tf.image.decode_png(img, channels=3) img = tf.cast(img, tf.float32) / 255.0 # Normalize to [0, 1] range # Load and preprocess segmentation mask mask = tf.io.read_file(mask_path) mask = tf.image.decode_png(mask, channels=1) # Replace ignore labels with a value your loss function will skip (e.g., -1) # Adjust `ignore_label_value` to match your dataset's ignore code (usually 255) mask = tf.where(mask == 255, tf.constant(-1, dtype=tf.int32), mask) return img, mask
Map this function to your dataset to load data as needed:
tr_data = tr_data.map(parse_image_mask, num_parallel_calls=tf.data.AUTOTUNE)
Using shuffle(len(train_in_np)) sets the shuffle buffer to the entire dataset size, which eats up unnecessary memory. Instead, use a reasonable buffer size (like 1000) to keep shuffling efficient:
tr_data = tr_data.shuffle(buffer_size=1000)
Semantic segmentation requires that images and masks are augmented in the same way (e.g., flipping both horizontally). Add an augmentation function that transforms both inputs together:
def augment_pair(image, mask): # Random horizontal flip if tf.random.uniform(()) > 0.5: image = tf.image.flip_left_right(image) mask = tf.image.flip_left_right(mask) # Add other augmentations (rotation, zoom, etc.) here if needed return image, mask
Then chain in batching and prefetching to optimize pipeline performance:
tr_data = tr_data.map(augment_pair, num_parallel_calls=tf.data.AUTOTUNE) tr_data = tr_data.batch(batch_size=8) # Adjust based on your GPU memory tr_data = tr_data.prefetch(tf.data.AUTOTUNE) # Overlap data preprocessing and model training
Hardcoding paths and training-only logic makes your function inflexible. Refactor it to accept a mode (train/eval) and other parameters, so you can use the same function for both training and validation:
def input_fn(mode, batch_size=8, basepath="./"): # Set paths based on mode if mode == tf.estimator.ModeKeys.TRAIN: img_dir = join(basepath, "leftImg8bit/train/*/*") mask_dir = join(basepath, "gtFine/train/*/*_ignoreLabel.png") shuffle = True else: img_dir = join(basepath, "leftImg8bit/val/*/*") mask_dir = join(basepath, "gtFine/val/*/*_ignoreLabel.png") shuffle = False # Load and sort paths to keep image-mask pairs matched img_paths = sorted(io_utils.loadDataset(img_dir)) mask_paths = sorted(io_utils.loadDataset(mask_dir)) # Create dataset dataset = tf.data.Dataset.from_tensor_slices((img_paths, mask_paths)) if shuffle: dataset = dataset.shuffle(buffer_size=1000) # Parse, augment, batch dataset = dataset.map(parse_image_mask, num_parallel_calls=tf.data.AUTOTUNE) if mode == tf.estimator.ModeKeys.TRAIN: dataset = dataset.map(augment_pair, num_parallel_calls=tf.data.AUTOTUNE) dataset = dataset.batch(batch_size) dataset = dataset.prefetch(tf.data.AUTOTUNE) return dataset
Make sure your model's loss function ignores the special value you assigned to ignore labels. For example, if you used -1, use SparseCategoricalCrossentropy with ignore_class=-1:
loss_fn = tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True, ignore_class=-1)
内容的提问来源于stack exchange,提问作者Simon Rühle

