File size: 1,751 Bytes
07cb7d3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
"""
图片分类模型构建组件,小型 Xception 风格二分类网络。
"""

from dataclasses import dataclass

import keras
from keras.layers import BatchNormalization, Conv2D, Dense, Dropout, GlobalAveragePooling2D, MaxPooling2D, Rescaling, SeparableConv2D

from deep_learning.models.spec import ModelArtifact, SupervisedModelBuilder


@dataclass
class ImageClassificationModelBuilder(SupervisedModelBuilder):
    image_size: tuple[int, int]
    model_filters: tuple[int, ...]

    def build_training_artifact(self) -> ModelArtifact:
        inputs = keras.Input(shape=self.image_size + (3,))
        x = Rescaling(1.0 / 255)(inputs)
        x = Conv2D(32, 3, strides=2, padding="same", use_bias=False)(x)

        for filter_count in self.model_filters:
            residual = Conv2D(filter_count, 1, strides=2, padding="same", use_bias=False)(x)
            residual = BatchNormalization()(residual)

            x = SeparableConv2D(filter_count, 3, padding="same", use_bias=False)(x)
            x = BatchNormalization()(x)
            x = keras.activations.relu(x)
            x = SeparableConv2D(filter_count, 3, padding="same", use_bias=False)(x)
            x = BatchNormalization()(x)
            x = MaxPooling2D(3, strides=2, padding="same")(x)
            x = keras.layers.add([x, residual])

        x = GlobalAveragePooling2D()(x)
        x = Dropout(0.5)(x)
        outputs = Dense(1, activation="sigmoid")(x)
        model = keras.Model(inputs, outputs, name="image_classification")
        return ModelArtifact(model=model)

    def compile_training_model(self, model: keras.Model) -> None:
        model.compile(
            optimizer="adam",
            loss="binary_crossentropy",
            metrics=["accuracy"]
        )