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

如何用Keras中的CNN处理马萨诸塞道路数据集并提取道路像素?

Extracting Road Pixels & Building a Keras CNN for Massachusetts Roads Dataset

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.20 10:40:40