Safe lab preview
Semantic Segmentation TF
This is a sanitized, read-only preview. Nothing executes in this page.
Read-only
Notebook preview
Semantic Segmentation TF
> **Integrated runtime note:** This preview uses a small deterministic, rights-safe offline fixture for reproducible learning. Full-scale results require the lesson's documented dataset or model in an approved external environment.
# Semantic Segmentation
**Segmentation** is one of the main computer vision task. For **each** pixel of image you must specify class(background included). Semantic segmentation only tells pixel class, instance segmentation divide classes into different instances.
For instance segmentation ten cars is **different** objects, for semantic segmentation **all** cars is one class.
> **Figure description:** The original was excluded because its reuse rights are not documented.
> Image from [this blog post](https://www.tensorflow.org/tutorials/images/segmentation)
Almost all architectures have same structure. First part is **encoder** that extracts features from input image, second part is **decoder** that transforms this features into image with same height and width and some number of channels, may be equal to classes count.
> **Figure description:** The original was excluded because its reuse rights are not documented.
> Image from [this publication](https://arxiv.org/pdf/2001.05566.pdf)
import tensorflow as tf
import tensorflow.keras.layers as keras
import tensorflow.keras.optimizers as optimizers
import tensorflow.keras.losses as losses
import matplotlib.pyplot as plt
import numpy as np
tf.random.set_seed(42)
np.random.seed(42)train_size = 0.8
lr = 3e-4
weight_decay = 8e-9
batch_size = 64
epochs = 2## Synthetic segmentation fixture
We generate edition-authored shapes and pixel masks so the segmentation example runs without downloading medical data or accepting an external dataset license.
# No package installation or external archive is required.def make_segmentation_dataset(samples=160, size=64, seed=2026):
rng = np.random.default_rng(seed)
y, x = np.mgrid[0:size, 0:size]
images, masks = [], []
for _ in range(samples):
cx, cy = rng.integers(size // 4, 3 * size // 4, size=2)
radius = int(rng.integers(size // 8, size // 4))
mask = ((x - cx) ** 2 + (y - cy) ** 2 <= radius ** 2).astype("float32")
image = rng.normal(0.12, 0.035, (size, size, 3)).astype("float32")
color = rng.uniform(0.55, 0.95, 3).astype("float32")
image[mask.astype(bool)] = color
images.append(np.clip(image, 0, 1))
masks.append(mask[..., None])
images = np.asarray(images, dtype="float32")
masks = np.asarray(masks, dtype="float32")
split = int(len(images) * train_size)
return (images[:split], masks[:split]), (images[split:], masks[split:])(X_train, y_train), (X_test, y_test) = make_segmentation_dataset()def plotn(n, data):
images, masks = data[0], data[1]
fig, ax = plt.subplots(1, n)
fig1, ax1 = plt.subplots(1, n)
for i, (img, mask) in enumerate(zip(images, masks)):
if i == n:
break
ax[i].imshow(img)
ax1[i].imshow(mask[:, :, 0])
plt.show()**Let's plot some images with corresponding masks.**
plotn(5, (X_train, y_train))## SegNet
Simple encoder - decoder architecture with convolutions, poolings in encoder and convolutions, upsamplings in decoder.
> **Figure description:** The original was excluded because its reuse rights are not documented.
* Badrinarayanan, V., Kendall, A., & Cipolla, R. (2015). [SegNet: A deep convolutional
encoder-decoder architecture for image segmentation](https://arxiv.org/pdf/1511.00561.pdf)
class SegNet(tf.keras.Model):
def __init__(self):
super().__init__()
self.enc_conv0 = keras.Conv2D(16, kernel_size=3, padding='same')
self.bn0 = keras.BatchNormalization()
self.relu0 = keras.Activation('relu')
self.pool0 = keras.MaxPool2D()
self.enc_conv1 = keras.Conv2D(32, kernel_size=3, padding='same')
self.relu1 = keras.Activation('relu')
self.bn1 = keras.BatchNormalization()
self.pool1 = keras.MaxPool2D()
self.enc_conv2 = keras.Conv2D(64, kernel_size=3, padding='same')
self.relu2 = keras.Activation('relu')
self.bn2 = keras.BatchNormalization()
self.pool2 = keras.MaxPool2D()
self.enc_conv3 = keras.Conv2D(128, kernel_size=3, padding='same')
self.relu3 = keras.Activation('relu')
self.bn3 = keras.BatchNormalization()
self.pool3 = keras.MaxPool2D()
self.bottleneck_conv = keras.Conv2D(256, kernel_size=(3, 3), padding='same')
self.upsample0 = keras.UpSampling2D(interpolation='bilinear')
self.dec_conv0 = keras.Conv2D(128, kernel_size=3, padding='same')
self.dec_relu0 = keras.Activation('relu')
self.dec_bn0 = keras.BatchNormalization()
self.upsample1 = keras.UpSampling2D(interpolation='bilinear')
self.dec_conv1 = keras.Conv2D(64, kernel_size=3, padding='same')
self.dec_relu1 = keras.Activation('relu')
self.dec_bn1 = keras.BatchNormalization()
self.upsample2 = keras.UpSampling2D(interpolation='bilinear')
self.dec_conv2 = keras.Conv2D(32, kernel_size=3, padding='same')
self.dec_relu2 = keras.Activation('relu')
self.dec_bn2 = keras.BatchNormalization()
self.upsample3 = keras.UpSampling2D(interpolation='bilinear')
self.dec_conv3 = keras.Conv2D(1, kernel_size=1)
def call(self, input):
e0 = self.pool0(self.relu0(self.bn0(self.enc_conv0(input))))
e1 = self.pool1(self.relu1(self.bn1(self.enc_conv1(e0))))
e2 = self.pool2(self.relu2(self.bn2(self.enc_conv2(e1))))
e3 = self.pool3(self.relu3(self.bn3(self.enc_conv3(e2))))
b = self.bottleneck_conv(e3)
d0 = self.dec_relu0(self.dec_bn0(self.upsample0(self.dec_conv0(b))))
d1 = self.dec_relu1(self.dec_bn1(self.upsample1(self.dec_conv1(d0))))
d2 = self.dec_relu2(self.dec_bn2(self.upsample2(self.dec_conv2(d1))))
d3 = self.dec_conv3(self.upsample3(d2))
return d3model = SegNet()
optimizer = optimizers.AdamW(learning_rate=lr, weight_decay=weight_decay)
loss_fn = losses.BinaryCrossentropy(from_logits=True)
model.compile(loss=loss_fn, optimizer=optimizer)def train(datasets, model, epochs, batch_size):
train_dataset, test_dataset = datasets[0], datasets[1]
model.fit(train_dataset[0], train_dataset[1],
epochs=epochs,
batch_size=batch_size,
shuffle=True,
validation_data=(test_dataset[0], test_dataset[1]))train(((X_train, y_train), (X_test, y_test)), model, epochs, batch_size)predictions = []
image_mask = []
plots = 5
for i, (img, mask) in enumerate(zip(X_test, y_test)):
if i == plots:
break
img = tf.expand_dims(img, 0)
pred = np.array(model.predict(img))
predictions.append(pred[0, :, :, 0] > 0.0)
image_mask.append(mask)
plotn(plots, (predictions, image_mask))## U-Net
Very simple architecture that uses skip connections. Skip connections at each convolution level helps network doesn't lost information about features from original input at this level.
U-Net usually has a default encoder for feature extraction, for example resnet50.
> **Figure description:** The original was excluded because its reuse rights are not documented.
* Ronneberger, Olaf, Philipp Fischer, and Thomas Brox. [U-Net: Convolutional networks for biomedical image segmentation.](https://arxiv.org/pdf/1505.04597.pdf)
class UNet(tf.keras.Model):
def __init__(self):
super().__init__()
self.enc_conv0 = keras.Conv2D(16, kernel_size=3, padding='same')
self.bn0 = keras.BatchNormalization()
self.relu0 = keras.Activation('relu')
self.pool0 = keras.MaxPool2D()
self.enc_conv1 = keras.Conv2D(32, kernel_size=3, padding='same')
self.relu1 = keras.Activation('relu')
self.bn1 = keras.BatchNormalization()
self.pool1 = keras.MaxPool2D()
self.enc_conv2 = keras.Conv2D(64, kernel_size=3, padding='same')
self.relu2 = keras.Activation('relu')
self.bn2 = keras.BatchNormalization()
self.pool2 = keras.MaxPool2D()
self.enc_conv3 = keras.Conv2D(128, kernel_size=3, padding='same')
self.relu3 = keras.Activation('relu')
self.bn3 = keras.BatchNormalization()
self.pool3 = keras.MaxPool2D()
self.bottleneck_conv = keras.Conv2D(256, kernel_size=(3, 3), padding='same')
self.upsample0 = keras.UpSampling2D(interpolation='bilinear')
self.dec_conv0 = keras.Conv2D(128, kernel_size=3, padding='same', input_shape=[None, 384, None, None])
self.dec_relu0 = keras.Activation('relu')
self.dec_bn0 = keras.BatchNormalization()
self.upsample1 = keras.UpSampling2D(interpolation='bilinear')
self.dec_conv1 = keras.Conv2D(64, kernel_size=3, padding='same', input_shape=[None, 192, None, None])
self.dec_relu1 = keras.Activation('relu')
self.dec_bn1 = keras.BatchNormalization()
self.upsample2 = keras.UpSampling2D(interpolation='bilinear')
self.dec_conv2 = keras.Conv2D(32, kernel_size=3, padding='same', input_shape=[None, 96, None, None])
self.dec_relu2 = keras.Activation('relu')
self.dec_bn2 = keras.BatchNormalization()
self.upsample3 = keras.UpSampling2D(interpolation='bilinear')
self.dec_conv3 = keras.Conv2D(1, kernel_size=1, input_shape=[None, 48, None, None])
self.cat0 = keras.Concatenate(axis=3)
self.cat1 = keras.Concatenate(axis=3)
self.cat2 = keras.Concatenate(axis=3)
self.cat3 = keras.Concatenate(axis=3)
def call(self, input):
e0 = self.pool0(self.relu0(self.bn0(self.enc_conv0(input))))
e1 = self.pool1(self.relu1(self.bn1(self.enc_conv1(e0))))
e2 = self.pool2(self.relu2(self.bn2(self.enc_conv2(e1))))
e3 = self.pool3(self.relu3(self.bn3(self.enc_conv3(e2))))
cat0 = self.relu0(self.bn0(self.enc_conv0(input)))
cat1 = self.relu1(self.bn1(self.enc_conv1(e0)))
cat2 = self.relu2(self.bn2(self.enc_conv2(e1)))
cat3 = self.relu3(self.bn3(self.enc_conv3(e2)))
b = self.bottleneck_conv(e3)
cat_tens0 = self.cat0([self.upsample0(b), cat3])
d0 = self.dec_relu0(self.dec_bn0(self.dec_conv0(cat_tens0)))
cat_tens1 = self.cat1([self.upsample1(d0), cat2])
d1 = self.dec_relu1(self.dec_bn1(self.dec_conv1(cat_tens1)))
cat_tens2 = self.cat2([self.upsample2(d1), cat1])
d2 = self.dec_relu2(self.dec_bn2(self.dec_conv2(cat_tens2)))
cat_tens3 = self.cat3([self.upsample3(d2), cat0])
d3 = self.dec_conv3(cat_tens3)
return d3model = UNet()
optimizer = optimizers.AdamW(learning_rate=lr, weight_decay=weight_decay)
loss_fn = losses.BinaryCrossentropy(from_logits=True)
model.compile(loss=loss_fn, optimizer=optimizer)train(((X_train, y_train), (X_test, y_test)), model, epochs, batch_size)predictions = []
image_mask = []
plots = 5
for i, (img, mask) in enumerate(zip(X_test, y_test)):
if i == plots:
break
img = tf.expand_dims(img, 0)
pred = np.array(model.predict(img))
predictions.append(pred[0, :, :, 0] > 0.0)
image_mask.append(mask)
plotn(plots, (predictions, image_mask))Outputs, execution counts, widgets, and active content were removed during import. Run notebooks only in an external environment you trust.
Record your practice
Optional self-reporting helps you remember what you practiced and never gates course completion.