Implementing Semantic Image Segmentation with FCN using MindSpore

Implementing Semantic Image Segmentation with FCN using MindSpore

Background

This article explores the implementation of Fully Convolutional Networks (FCN) for semantic image segmentation using MindSpore, an open-source deep learning framework developed by Huawei.

FCN Introduction

Fully Convolutional Networks (FCN) represent a significant advancement in semantic image segmentation, first introduced by Jonathan Long and colleagues from UC Berkeley in 2015. The FCN framework was groundbreaking as it was the first end-to-end approach capable of pixel-level predictions.

Key innovations of FCN include:

  1. Convolutional Architecture: Utilizing VGG-16 as a backbone network with fully connected layers replaced by convolutional layers, transforming one-dimensional outputs into two-dimensional heatmaps.
  2. Upsampling Operations: Addressing the reduction in feature map dimensions caused by convolution and pooling operations through upsampling to restore original image sizes.
  3. Skip Architecture: Combining predictions from deeper layers with more global information and shallower layers with finer details to enhance segmentation accuracy.

Characteristics of the FCN architecture:

  • Fully convolutional without fully connected layers, enabling input of arbitrary dimensions
  • Transposed convolution layers to increase spatial dimensions
  • Skip connections combining features from different network depths for improved robustness and precision

Implementation

Data Download

from download import download
dataset_url = "https://mindspore-website.obs.cn-north-4.myhuaweicloud.com/notebook/datasets/dataset_fcn8s.tar"
download(dataset_url, "./dataset", kind="tar", replace=True)

Data Preprocessing

import numpy as np
import cv2
import mindspore.dataset as ds

class SemanticSegmentationDataset:
    def __init__(self,
                 image_mean_values,
                 image_std_values,
                 data_path='',
                 batch_size=32,
                 crop_dimension=512,
                 max_scaling_factor=2.0,
                 min_scaling_factor=0.5,
                 ignored_label=255,
                 num_categories=21,
                 num_data_readers=2,
                 parallel_workers=4):

        self.data_path = data_path
        self.batch_size = batch_size
        self.crop_dimension = crop_dimension
        self.image_mean = np.array(image_mean_values, dtype=np.float32)
        self.image_std = np.array(image_std_values, dtype=np.float32)
        self.max_scaling = max_scaling_factor
        self.min_scaling = min_scaling_factor
        self.ignored_label = ignored_label
        self.num_categories = num_categories
        self.num_readers = num_data_readers
        self.parallel_workers = parallel_workers
        assert max_scaling_factor > min_scaling_factor

    def prepare_data(self, image_data, label_data):
        image_processed = cv2.imdecode(np.frombuffer(image_data, dtype=np.uint8), cv2.IMREAD_COLOR)
        label_processed = cv2.imdecode(np.frombuffer(label_data, dtype=np.uint8), cv2.IMREAD_GRAYSCALE)
        
        scaling_factor = np.random.uniform(self.min_scaling, self.max_scaling)
        new_height, new_width = int(scaling_factor * image_processed.shape[0]), int(scaling_factor * image_processed.shape[1])
        
        image_processed = cv2.resize(image_processed, (new_width, new_height), interpolation=cv2.INTER_CUBIC)
        label_processed = cv2.resize(label_processed, (new_width, new_height), interpolation=cv2.INTER_NEAREST)

        image_processed = (image_processed - self.image_mean) / self.image_std
        
        output_height, output_width = max(new_height, self.crop_dimension), max(new_width, self.crop_dimension)
        padding_height, padding_width = output_height - new_height, output_width - new_width
        
        if padding_height > 0 or padding_width > 0:
            image_processed = cv2.copyMakeBorder(image_processed, 0, padding_height, 0, padding_width, cv2.BORDER_CONSTANT, value=0)
            label_processed = cv2.copyMakeBorder(label_processed, 0, padding_height, 0, padding_width, cv2.BORDER_CONSTANT, value=self.ignored_label)
            
        height_offset = np.random.randint(0, output_height - self.crop_dimension + 1)
        width_offset = np.random.randint(0, output_width - self.crop_dimension + 1)
        
        image_processed = image_processed[height_offset: height_offset + self.crop_dimension, 
                                         width_offset: width_offset + self.crop_dimension, :]
        label_processed = label_processed[height_offset: height_offset + self.crop_dimension, 
                                          width_offset: width_offset + self.crop_dimension]
        
        if np.random.uniform(0.0, 1.0) > 0.5:
            image_processed = image_processed[:, ::-1, :]
            label_processed = label_processed[:, ::-1]
            
        image_processed = image_processed.transpose((2, 0, 1))
        image_processed = image_processed.copy()
        label_processed = label_processed.copy()
        label_processed = label_processed.astype("int32")
        
        return image_processed, label_processed

    def load_dataset(self):
        ds.config.set_numa_enable(True)
        dataset = ds.MindDataset(self.data_path, columns_list=["image", "mask"],
                                 shuffle=True, num_parallel_workers=self.num_readers)
        
        transformations = self.prepare_data
        dataset = dataset.map(operations=transformations, input_columns=["image", "mask"],
                              output_columns=["image", "mask"],
                              num_parallel_workers=self.parallel_workers)
        
        dataset = dataset.shuffle(buffer_size=self.batch_size * 10)
        dataset = dataset.batch(self.batch_size, drop_remainder=True)
        
        return dataset


# Configuration parameters for dataset creation
IMAGE_AVERAGES = [103.53, 116.28, 123.675]
IMAGE_DEVIATIONS = [57.375, 57.120, 58.395]
DATA_FILE_PATH = "dataset/dataset_fcn8s/mindname.mindrecord"

# Model training configuration
training_batch_size = 4
crop_size = 512
minimum_scale = 0.5
maximum_scale = 2.0
ignored_label = 255
num_classes = 21

# Instantiate the dataset
dataset_handler = SemanticSegmentationDataset(
    image_mean_values=IMAGE_AVERAGES,
    image_std_values=IMAGE_DEVIATIONS,
    data_path=DATA_FILE_PATH,
    batch_size=training_batch_size,
    crop_dimension=crop_size,
    max_scaling_factor=maximum_scale,
    min_scaling_factor=minimum_scale,
    ignored_label=ignored_label,
    num_categories=num_classes,
    num_data_readers=2,
    parallel_workers=4
)

dataset = dataset_handler.load_dataset()

Model Architecture

import mindspore.nn as nn

class FCN8sNetwork(nn.Cell):
    def __init__(self, num_output_classes):
        super().__init__()
        self.num_classes = num_output_classes
        
        # Initial convolutional blocks
        self.initial_conv = nn.SequentialCell(
            nn.Conv2d(in_channels=3, out_channels=64,
                      kernel_size=3, weight_init='xavier_uniform'),
            nn.BatchNorm2d(64),
            nn.ReLU(),
            nn.Conv2d(in_channels=64, out_channels=64,
                      kernel_size=3, weight_init='xavier_uniform'),
            nn.BatchNorm2d(64),
            nn.ReLU()
        )
        self.initial_pool = nn.MaxPool2d(kernel_size=2, stride=2)
        
        # Second convolutional block
        self.second_conv = nn.SequentialCell(
            nn.Conv2d(in_channels=64, out_channels=128,
                      kernel_size=3, weight_init='xavier_uniform'),
            nn.BatchNorm2d(128),
            nn.ReLU(),
            nn.Conv2d(in_channels=128, out_channels=128,
                      kernel_size=3, weight_init='xavier_uniform'),
            nn.BatchNorm2d(128),
            nn.ReLU()
        )
        self.second_pool = nn.MaxPool2d(kernel_size=2, stride=2)
        
        # Third convolutional block
        self.third_conv = nn.SequentialCell(
            nn.Conv2d(in_channels=128, out_channels=256,
                      kernel_size=3, weight_init='xavier_uniform'),
            nn.BatchNorm2d(256),
            nn.ReLU(),
            nn.Conv2d(in_channels=256, out_channels=256,
                      kernel_size=3, weight_init='xavier_uniform'),
            nn.BatchNorm2d(256),
            nn.ReLU(),
            nn.Conv2d(in_channels=256, out_channels=256,
                      kernel_size=3, weight_init='xavier_uniform'),
            nn.BatchNorm2d(256),
            nn.ReLU()
        )
        self.third_pool = nn.MaxPool2d(kernel_size=2, stride=2)
        
        # Fourth convolutional block
        self.fourth_conv = nn.SequentialCell(
            nn.Conv2d(in_channels=256, out_channels=512,
                      kernel_size=3, weight_init='xavier_uniform'),
            nn.BatchNorm2d(512),
            nn.ReLU(),
            nn.Conv2d(in_channels=512, out_channels=512,
                      kernel_size=3, weight_init='xavier_uniform'),
            nn.BatchNorm2d(512),
            nn.ReLU(),
            nn.Conv2d(in_channels=512, out_channels=512,
                      kernel_size=3, weight_init='xavier_uniform'),
            nn.BatchNorm2d(512),
            nn.ReLU()
        )
        self.fourth_pool = nn.MaxPool2d(kernel_size=2, stride=2)
        
        # Fifth convolutional block
        self.fifth_conv = nn.SequentialCell(
            nn.Conv2d(in_channels=512, out_channels=512,
                      kernel_size=3, weight_init='xavier_uniform'),
            nn.BatchNorm2d(512),
            nn.ReLU(),
            nn.Conv2d(in_channels=512, out_channels=512,
                      kernel_size=3, weight_init='xavier_uniform'),
            nn.BatchNorm2d(512),
            nn.ReLU(),
            nn.Conv2d(in_channels=512, out_channels=512,
                      kernel_size=3, weight_init='xavier_uniform'),
            nn.BatchNorm2d(512),
            nn.ReLU()
        )
        self.fifth_pool = nn.MaxPool2d(kernel_size=2, stride=2)
        
        # Final convolutional layers
        self.sixth_conv = nn.SequentialCell(
            nn.Conv2d(in_channels=512, out_channels=4096,
                      kernel_size=7, weight_init='xavier_uniform'),
            nn.BatchNorm2d(4096),
            nn.ReLU(),
        )
        self.seventh_conv = nn.SequentialCell(
            nn.Conv2d(in_channels=4096, out_channels=4096,
                      kernel_size=1, weight_init='xavier_uniform'),
            nn.BatchNorm2d(4096),
            nn.ReLU(),
        )
        
        # Scoring layers
        self.feature_scoring = nn.Conv2d(in_channels=4096, out_channels=self.num_classes,
                                         kernel_size=1, weight_init='xavier_uniform')
        
        # Upsampling layers
        self.double_upsampling = nn.Conv2dTranspose(in_channels=self.num_classes, out_channels=self.num_classes,
                                                   kernel_size=4, stride=2, weight_init='xavier_uniform')
        self.pool4_scoring = nn.Conv2d(in_channels=512, out_channels=self.num_classes,
                                      kernel_size=1, weight_init='xavier_uniform')
        self.pool4_upsampling = nn.Conv2dTranspose(in_channels=self.num_classes, out_channels=self.num_classes,
                                                   kernel_size=4, stride=2, weight_init='xavier_uniform')
        self.pool3_scoring = nn.Conv2d(in_channels=256, out_channels=self.num_classes,
                                      kernel_size=1, weight_init='xavier_uniform')
        self.eight_upsampling = nn.Conv2dTranspose(in_channels=self.num_classes, out_channels=self.num_classes,
                                                   kernel_size=16, stride=8, weight_init='xavier_uniform')

    def forward(self, input_data):
        # Encoder path
        x1 = self.initial_conv(input_data)
        p1 = self.initial_pool(x1)
        x2 = self.second_conv(p1)
        p2 = self.second_pool(x2)
        x3 = self.third_conv(p2)
        p3 = self.third_pool(x3)
        x4 = self.fourth_conv(p3)
        p4 = self.fourth_pool(x4)
        x5 = self.fifth_conv(p4)
        p5 = self.fifth_pool(x5)
        
        # Decoder path with skip connections
        x6 = self.sixth_conv(p5)
        x7 = self.seventh_conv(x6)
        score_features = self.feature_scoring(x7)
        
        # Skip connections
        upsampled_2x = self.double_upsampling(score_features)
        pool4_score = self.pool4_scoring(p4)
        fused_pool4 = pool4_score + upsampled_2x
        upsampled_pool4 = self.pool4_upsampling(fused_pool4)
        pool3_score = self.pool3_scoring(p3)
        fused_pool3 = pool3_score + upsampled_pool4
        final_output = self.eight_upsampling(fused_pool3)
        
        return final_output

Loss Function and Evaluation Metrics

import numpy as np
import mindspore as ms
import mindspore.nn as nn
import mindspore.train as train

class SegmentationAccuracy(train.Metric):
    def __init__(self, num_categories=21):
        super(SegmentationAccuracy, self).__init__()
        self.num_categories = num_categories

    def _create_confusion_matrix(self, ground_truth, predicted):
        valid_mask = (ground_truth >= 0) & (ground_truth < self.num_categories)
        label_indices = self.num_categories * ground_truth[valid_mask].astype('int') + predicted[valid_mask]
        bin_counts = np.bincount(label_indices, minlength=self.num_categories**2)
        confusion_matrix = bin_counts.reshape(self.num_categories, self.num_categories)
        return confusion_matrix

    def clear(self):
        self.confusion_matrix = np.zeros((self.num_categories,) * 2)

    def update(self, *inputs):
        predictions = inputs[0].asnumpy().argmax(axis=1)
        ground_truth = inputs[1].asnumpy().reshape(4, 512, 512)
        self.confusion_matrix += self._create_confusion_matrix(ground_truth, predictions)

    def evaluate(self):
        pixel_accuracy = np.diag(self.confusion_matrix).sum() / self.confusion_matrix.sum()
        return pixel_accuracy


class CategoryAccuracy(train.Metric):
    def __init__(self, num_categories=21):
        super(CategoryAccuracy, self).__init__()
        self.num_categories = num_categories

    def _create_confusion_matrix(self, ground_truth, predicted):
        valid_mask = (ground_truth >= 0) & (ground_truth < self.num_categories)
        label_indices = self.num_categories * ground_truth[valid_mask].astype('int') + predicted[valid_mask]
        bin_counts = np.bincount(label_indices, minlength=self.num_categories**2)
        confusion_matrix = bin_counts.reshape(self.num_categories, self.num_categories)
        return confusion_matrix

    def update(self, *inputs):
        predictions = inputs[0].asnumpy().argmax(axis=1)
        ground_truth = inputs[1].asnumpy().reshape(4, 512, 512)
        self.confusion_matrix += self._create_confusion_matrix(ground_truth, predictions)

    def clear(self):
        self.confusion_matrix = np.zeros((self.num_categories,) * 2)

    def evaluate(self):
        category_accuracies = np.diag(self.confusion_matrix) / self.confusion_matrix.sum(axis=1)
        mean_category_accuracy = np.nanmean(category_accuracies)
        return mean_category_accuracy


class MeanIoU(train.Metric):
    def __init__(self, num_categories=21):
        super(MeanIoU, self).__init__()
        self.num_categories = num_categories

    def _create_confusion_matrix(self, ground_truth, predicted):
        valid_mask = (ground_truth >= 0) & (ground_truth < self.num_categories)
        label_indices = self.num_categories * ground_truth[valid_mask].astype('int') + predicted[valid_mask]
        bin_counts = np.bincount(label_indices, minlength=self.num_categories**2)
        confusion_matrix = bin_counts.reshape(self.num_categories, self.num_categories)
        return confusion_matrix

    def update(self, *inputs):
        predictions = inputs[0].asnumpy().argmax(axis=1)
        ground_truth = inputs[1].asnumpy().reshape(4, 512, 512)
        self.confusion_matrix += self._create_confusion_matrix(ground_truth, predictions)

    def clear(self):
        self.confusion_matrix = np.zeros((self.num_categories,) * 2)

    def evaluate(self):
        intersection_over_union = np.diag(self.confusion_matrix) / (
            np.sum(self.confusion_matrix, axis=1) + np.sum(self.confusion_matrix, axis=0) -
            np.diag(self.confusion_matrix))
        mean_iou = np.nanmean(intersection_over_union)
        return mean_iou


class FrequencyWeightedIoU(train.Metric):
    def __init__(self, num_categories=21):
        super(FrequencyWeightedIoU, self).__init__()
        self.num_categories = num_categories

    def _create_confusion_matrix(self, ground_truth, predicted):
        valid_mask = (ground_truth >= 0) & (ground_truth < self.num_categories)
        label_indices = self.num_categories * ground_truth[valid_mask].astype('int') + predicted[valid_mask]
        bin_counts = np.bincount(label_indices, minlength=self.num_categories**2)
        confusion_matrix = bin_counts.reshape(self.num_categories, self.num_categories)
        return confusion_matrix

    def update(self, *inputs):
        predictions = inputs[0].asnumpy().argmax(axis=1)
        ground_truth = inputs[1].asnumpy().reshape(4, 512, 512)
        self.confusion_matrix += self._create_confusion_matrix(ground_truth, predictions)

    def clear(self):
        self.confusion_matrix = np.zeros((self.num_categories,) * 2)

    def evaluate(self):
        category_frequencies = np.sum(self.confusion_matrix, axis=1) / np.sum(self.confusion_matrix)
        iou_scores = np.diag(self.confusion_matrix) / (
            np.sum(self.confusion_matrix, axis=1) + np.sum(self.confusion_matrix, axis=0) -
            np.diag(self.confusion_matrix))

        frequency_weighted_iou = (category_frequencies[category_frequencies > 0] * iou_scores[category_frequencies > 0]).sum()
        return frequency_weighted_iou

Model Training

import mindspore
from mindspore import Tensor
import mindspore.nn as nn
from mindspore.train import ModelCheckpoint, CheckpointConfig, LossMonitor, TimeMonitor, Model

# Set device context
device_platform = "Ascend"
mindspore.set_context(mode=mindspore.PYNATIVE_MODE, device_target=device_platform)

# Training configuration
batch_size = 4
num_classes = 21

# Initialize model architecture
segmentation_model = FCN8sNetwork(num_output_classes=21)
# Load VGG-16 pre-trained weights
load_vgg16()

# Learning rate scheduling
min_learning_rate = 0.0005
base_learning_rate = 0.05
training_epochs = 1
iterations_per_epoch = dataset.get_dataset_size()
total_iterations = iterations_per_epoch * training_epochs

learning_rate_scheduler = mindspore.nn.cosine_decay_lr(min_learning_rate,
                                                       base_learning_rate,
                                                       total_iterations,
                                                       iterations_per_epoch,
                                                       decay_epoch=2)
final_lr = Tensor(learning_rate_scheduler[-1])

# Define loss function
segmentation_loss = nn.CrossEntropyLoss(ignore_index=255)
# Define optimizer
training_optimizer = nn.Momentum(params=segmentation_model.trainable_params(), 
                                learning_rate=final_lr, 
                                momentum=0.9, 
                                weight_decay=0.0001)

# Configure loss scaling for mixed precision training
scaling_factor = 4
scaling_window = 3000
loss_scale_manager = ms.amp.DynamicLossScaleManager(scaling_factor, scaling_window)

# Initialize model
if device_platform == "Ascend":
    training_model = Model(segmentation_model, 
                          loss_fn=segmentation_loss, 
                          optimizer=training_optimizer, 
                          loss_scale_manager=loss_scale_manager,
                          metrics={"pixel accuracy": SegmentationAccuracy(), 
                                  "mean pixel accuracy": CategoryAccuracy(), 
                                  "mean IoU": MeanIoU(), 
                                  "frequency weighted IoU": FrequencyWeightedIoU()})
else:
    training_model = Model(segmentation_model, 
                          loss_fn=segmentation_loss, 
                          optimizer=training_optimizer,
                          metrics={"pixel accuracy": SegmentationAccuracy(), 
                                  "mean pixel accuracy": CategoryAccuracy(), 
                                  "mean IoU": MeanIoU(), 
                                  "frequency weighted IoU": FrequencyWeightedIoU()})

# Configure checkpoint saving
time_callback = TimeMonitor(data_size=iterations_per_epoch)
loss_callback = LossMonitor()
training_callbacks = [time_callback, loss_callback]

checkpoint_interval = 330
max_checkpoints = 5
checkpoint_config = CheckpointConfig(save_checkpoint_steps=10,
                                     keep_checkpoint_max=max_checkpoints)
checkpoint_callback = ModelCheckpoint(prefix="FCN8s",
                                     directory="./model_checkpoints",
                                     config=checkpoint_config)
training_callbacks.append(checkpoint_callback)

# Start training
training_model.train(training_epochs, dataset, callbacks=training_callbacks)

Model Evaluation

# Dataset configuration for evaluation
IMAGE_AVERAGES = [103.53, 116.28, 123.675]
IMAGE_DEVIATIONS = [57.375, 57.120, 58.395]
DATA_FILE_PATH = "dataset/dataset_fcn8s/mindname.mindrecord"

# Download pre-trained weights
model_url = "https://mindspore-website.obs.cn-north-4.myhuaweicloud.com/notebook/datasets/FCN8s.ckpt"
download(model_url, "FCN8s.ckpt", replace=True)

# Initialize model with pre-trained weights
evaluation_model = FCN8sNetwork(num_output_classes=num_classes)

checkpoint_path = "FCN8s.ckpt"
model_parameters = load_checkpoint(checkpoint_path)
load_param_into_net(evaluation_model, model_parameters)

# Create evaluation model
if device_platform == "Ascend":
    evaluation_engine = Model(evaluation_model, 
                            loss_fn=segmentation_loss, 
                            optimizer=training_optimizer, 
                            loss_scale_manager=loss_scale_manager,
                            metrics={"pixel accuracy": SegmentationAccuracy(), 
                                    "mean pixel accuracy": CategoryAccuracy(), 
                                    "mean IoU": MeanIoU(), 
                                    "frequency weighted IoU": FrequencyWeightedIoU()})
else:
    evaluation_engine = Model(evaluation_model, 
                            loss_fn=segmentation_loss, 
                            optimizer=training_optimizer,
                            metrics={"pixel accuracy": SegmentationAccuracy(), 
                                    "mean pixel accuracy": CategoryAccuracy(), 
                                    "mean IoU": MeanIoU(), 
                                    "frequency weighted IoU": FrequencyWeightedIoU()})

# Create evaluation dataset
evaluation_dataset_handler = SemanticSegmentationDataset(
    image_mean_values=IMAGE_AVERAGES,
    image_std_values=IMAGE_DEVIATIONS,
    data_path=DATA_FILE_PATH,
    batch_size=batch_size,
    crop_dimension=crop_size,
    max_scaling_factor=maximum_scale,
    min_scaling_factor=minimum_scale,
    ignored_label=ignored_label,
    num_categories=num_classes,
    num_data_readers=2,
    parallel_workers=4
)
evaluation_dataset = evaluation_dataset_handler.load_dataset()

# Run evaluation
evaluation_engine.eval(evaluation_dataset)

Model Inference

import cv2
import matplotlib.pyplot as plt

# Load trained model
inference_model = FCN8sNetwork(num_output_classes=num_classes)
checkpoint_file = "FCN8s.ckpt"
model_weights = load_checkpoint(checkpoint_file)
load_param_into_net(inference_model, model_weights)

# Configure batch size for inference
inference_batch_size = 4

# Prepare to display results
input_images = []
ground_truth_masks = []
inference_results = []

# Create visualization grid
plt.figure(figsize=(8, 5))
sample_data = next(evaluation_dataset.create_dict_iterator())
sample_images = sample_data["image"].asnumpy()
mask_images = sample_data["label"].reshape([4, 512, 512])
sample_images = np.clip(sample_images, 0, 1)

# Collect sample images
for i in range(inference_batch_size):
    input_images.append(sample_images[i])
    ground_truth_masks.append(mask_images[i])

# Run inference
model_output = inference_model(sample_data["image"]).asnumpy().argmax(axis=1)

# Display results
for i in range(inference_batch_size):
    plt.subplot(2, 4, i + 1)
    plt.imshow(input_images[i].transpose(1, 2, 0))
    plt.axis("off")
    plt.subplots_adjust(wspace=0.05, hspace=0.02)
    
    plt.subplot(2, 4, i + 5)
    plt.imshow(model_output[i])
    plt.axis("off")
    plt.subplots_adjust(wspace=0.05, hspace=0.02)

plt.show()

Tags: mindspore FCN Semantic Segmentation Deep Learning Computer Vision

Posted on Tue, 06 Oct 2026 16:25:13 +0000 by apw