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:
- 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.
- Upsampling Operations: Addressing the reduction in feature map dimensions caused by convolution and pooling operations through upsampling to restore original image sizes.
- 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()