general-deep-learning / test /models /segmentation_model_test.py
yetrun's picture
ver3: 将源码迁入 src/deep_learning 包,重塑训练流水线,规范 data/model 契约
07cb7d3
Raw
History Blame Contribute Delete
1.33 kB
from unittest.mock import Mock
import numpy as np
import tensorflow as tf
from deep_learning.models.segmentation import SegmentationModelBuilder
def test_segmentation_model_builder_builds_pixel_classifier():
"""验证分割模型保持输入分辨率,并为每个像素输出类别概率。"""
artifact = SegmentationModelBuilder(
image_size=(32, 32),
num_classes=3,
model_filters=(8,)
).build_training_artifact()
model = artifact.model
images = tf.zeros((2, 32, 32, 3), dtype=tf.float32)
outputs = model(images)
assert outputs.shape == (2, 32, 32, 3)
np.testing.assert_allclose(
tf.reduce_sum(outputs, axis=-1).numpy(),
np.ones((2, 32, 32)),
atol=1e-5
)
def test_segmentation_model_builder_compiles_training_model():
"""验证分割模型构建器会使用稀疏多分类和前景 IoU 编译模型。"""
model = Mock()
builder = SegmentationModelBuilder(
image_size=(32, 32),
num_classes=3,
model_filters=(8,)
)
builder.compile_training_model(model)
model.compile.assert_called_once()
_, kwargs = model.compile.call_args
assert kwargs["optimizer"] == "adam"
assert kwargs["loss"] == "sparse_categorical_crossentropy"
assert kwargs["metrics"][0].name == "foreground_iou"