如何用Keras中的CNN处理马萨诸塞道路数据集并提取道路像素?
Hey there! As someone who's worked with CNNs for semantic segmentation, let's break down your question into two actionable parts. I'll keep things practical with code examples you can run right away.
Part 1: Extract & Display Road-Only Pixels from Input Image
The sample output you have is a binary mask—white pixels (value 255) represent roads, black (0) represents everything else. To extract roads from the satellite image, we just need to use this mask as a filter. Here's how to do it with Python's PIL and NumPy:
Step-by-Step Code
from PIL import Image import numpy as np import matplotlib.pyplot as plt # Load your downloaded images (replace paths with your local file paths) sat_image = Image.open("10078660_15.tiff") road_mask = Image.open("10078660_15.tif") # Convert images to NumPy arrays for pixel-level operations sat_array = np.array(sat_image) mask_array = np.array(road_mask) # Create a boolean mask where True = road pixels (mask value 255) road_filter = mask_array == 255 # Apply the filter to the satellite image: set non-road pixels to black road_only_pixels = sat_array.copy() road_only_pixels[~road_filter] = 0 # ~ inverts the boolean mask # Display the results plt.figure(figsize=(15, 5)) plt.subplot(131) plt.imshow(sat_image) plt.title("Original Satellite Image") plt.axis("off") plt.subplot(132) plt.imshow(road_mask, cmap="gray") plt.title("Road Mask") plt.axis("off") plt.subplot(133) plt.imshow(road_only_pixels) plt.title("Extracted Road Pixels") plt.axis("off") plt.show()
How It Works
- The mask acts as a "stencil": we only keep pixels in the satellite image where the mask is white.
- Setting non-road pixels to
0(black) makes the roads stand out clearly. You could also set them to transparent if you prefer—just use an RGBA image format.
Part 2: Keras CNN for Massachusetts Roads Dataset
This is a semantic segmentation task (classifying every pixel as road or non-road), not image classification. The go-to model for this kind of task is U-Net—it's designed for small datasets and excels at capturing fine-grained details like road edges.
Step 1: Data Preparation
First, organize your dataset like this (standard for segmentation tasks):
mass_roads/ ├── train/ │ ├── sat/ # Training satellite images │ └── map/ # Corresponding road masks └── val/ ├── sat/ # Validation satellite images └── map/ # Corresponding road masks
Step 2: Build the U-Net Model
U-Net has an encoder (downsamples to extract features) and a decoder (upsamples to reconstruct pixel-level predictions) with skip connections to preserve spatial details.
from keras.models import Model from keras.layers import Input, Conv2D, MaxPooling2D, UpSampling2D, concatenate from keras.optimizers import Adam from keras.callbacks import EarlyStopping, ModelCheckpoint def build_unet(input_shape=(256, 256, 3)): # Input layer inputs = Input(input_shape) # Encoder: Downsample to capture features c1 = Conv2D(64, (3, 3), activation='relu', padding='same', kernel_initializer='he_normal')(inputs) c1 = Conv2D(64, (3, 3), activation='relu', padding='same', kernel_initializer='he_normal')(c1) p1 = MaxPooling2D((2, 2))(c1) c2 = Conv2D(128, (3, 3), activation='relu', padding='same', kernel_initializer='he_normal')(p1) c2 = Conv2D(128, (3, 3), activation='relu', padding='same', kernel_initializer='he_normal')(c2) p2 = MaxPooling2D((2, 2))(c2) c3 = Conv2D(256, (3, 3), activation='relu', padding='same', kernel_initializer='he_normal')(p2) c3 = Conv2D(256, (3, 3), activation='relu', padding='same', kernel_initializer='he_normal')(c3) p3 = MaxPooling2D((2, 2))(c3) # Bottleneck: Highest level of feature abstraction c4 = Conv2D(512, (3, 3), activation='relu', padding='same', kernel_initializer='he_normal')(p3) c4 = Conv2D(512, (3, 3), activation='relu', padding='same', kernel_initializer='he_normal')(c4) # Decoder: Upsample and combine with encoder features (skip connections) u5 = UpSampling2D((2, 2))(c4) u5 = concatenate([u5, c3]) # Skip connection from encoder c5 = Conv2D(256, (3, 3), activation='relu', padding='same', kernel_initializer='he_normal')(u5) c5 = Conv2D(256, (3, 3), activation='relu', padding='same', kernel_initializer='he_normal')(c5) u6 = UpSampling2D((2, 2))(c5) u6 = concatenate([u6, c2]) c6 = Conv2D(128, (3, 3), activation='relu', padding='same', kernel_initializer='he_normal')(u6) c6 = Conv2D(128, (3, 3), activation='relu', padding='same', kernel_initializer='he_normal')(c6) u7 = UpSampling2D((2, 2))(c6) u7 = concatenate([u7, c1]) c7 = Conv2D(64, (3, 3), activation='relu', padding='same', kernel_initializer='he_normal')(u7) c7 = Conv2D(64, (3, 3), activation='relu', padding='same', kernel_initializer='he_normal')(c7) # Output layer: Binary prediction (road=1, non-road=0) outputs = Conv2D(1, (1, 1), activation='sigmoid')(c7) # Compile the model model = Model(inputs=[inputs], outputs=[outputs]) model.compile(optimizer=Adam(learning_rate=1e-4), loss='binary_crossentropy', metrics=['accuracy']) return model # Initialize the model model = build_unet() model.summary()
Step 3: Data Loading & Augmentation
Data augmentation is critical here (the dataset isn't huge) to prevent overfitting. We'll use ImageDataGenerator to generate augmented images on the fly:
from keras.preprocessing.image import ImageDataGenerator def create_data_generator(sat_dir, mask_dir, batch_size=8, img_size=(256, 256)): # Augmentation settings for satellite images sat_datagen = ImageDataGenerator( rescale=1./255, horizontal_flip=True, vertical_flip=True, rotation_range=15, zoom_range=0.1 ) # Augmentation settings for masks (same as satellite to keep alignment) mask_datagen = ImageDataGenerator( rescale=1./255, horizontal_flip=True, vertical_flip=True, rotation_range=15, zoom_range=0.1 ) # Generate satellite images sat_generator = sat_datagen.flow_from_directory( sat_dir, target_size=img_size, class_mode=None, batch_size=batch_size, seed=42 # Same seed to keep mask-image alignment ) # Generate masks (grayscale since they're binary) mask_generator = mask_datagen.flow_from_directory( mask_dir, target_size=img_size, class_mode=None, batch_size=batch_size, color_mode='grayscale', seed=42 ) return zip(sat_generator, mask_generator) # Create train and validation generators train_generator = create_data_generator("mass_roads/train/sat", "mass_roads/train/map") val_generator = create_data_generator("mass_roads/val/sat", "mass_roads/val/map")
Step 4: Train the Model
Use callbacks to save the best model and stop training if validation loss stops improving:
# Callbacks early_stop = EarlyStopping(patience=5, monitor='val_loss', restore_best_weights=True) checkpoint = ModelCheckpoint("best_road_segmentation_model.h5", monitor='val_loss', save_best_only=True) # Train the model history = model.fit( train_generator, steps_per_epoch=len(train_generator), epochs=50, validation_data=val_generator, validation_steps=len(val_generator), callbacks=[early_stop, checkpoint] )
Step 5: Predict & Visualize Results
Once trained, use the model to predict road masks on new images:
def predict_road_mask(image_path, model, img_size=(256, 256)): # Load and preprocess the image img = Image.open(image_path).resize(img_size) img_array = np.array(img) / 255.0 img_array = np.expand_dims(img_array, axis=0) # Add batch dimension # Predict the mask pred_mask = model.predict(img_array)[0] pred_mask = (pred_mask > 0.5).astype(np.uint8) # Threshold to binary (0/1) # Visualize plt.figure(figsize=(10, 5)) plt.subplot(121) plt.imshow(img) plt.title("Input Satellite Image") plt.axis("off") plt.subplot(122) plt.imshow(pred_mask, cmap="gray") plt.title("Predicted Road Mask") plt.axis("off") plt.show() # Test with a new image predict_road_mask("mass_roads/val/sat/test_image.tiff", model)
Key Tips for Success
- Dice Coefficient: For segmentation, accuracy isn't always the best metric. Add Dice coefficient as a custom metric to better evaluate performance (it measures overlap between predicted and true masks).
- Class Imbalance: Roads often cover a small portion of the image. Use weighted loss functions or oversample road pixels to balance the model's learning.
- Transfer Learning: If you have limited data, use a pre-trained encoder (like VGG16) instead of training from scratch—this can speed up convergence and improve results.
内容的提问来源于stack exchange,提问作者Faraz Gerrard Jamal

