File size: 3,088 Bytes
f5cf82d
 
f116465
f5cf82d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c01459e
f5cf82d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
from transformers import PreTrainedConfig

class NulaConfig(PreTrainedConfig):
    model_type="nula"
    def __init__(
        self, num_classes=10, in_channels=3, input_size=(3, 32, 32), block_channels=(128, 256, 512), 
        use_residual=True, use_se=True, use_spatial_attention=False, se_reduction=16, classifier_hidden_dim=256,
        mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5), id2label=None, label2id=None, **kwargs
    ):
        super().__init__(**kwargs)

        self.num_classes = num_classes
        self.in_channels = in_channels
        self.input_size = input_size
        self.block_channels = list(block_channels)
        
        self.use_residual = use_residual
        self.use_se = use_se
        self.use_spatial_attention = use_spatial_attention
        self.se_reduction = se_reduction
        
        self.classifier_hidden_dim = classifier_hidden_dim
        self.mean = list(mean)
        self.std = list(std)

        if id2label is None or label2id is None:
            labels = [
                "airplane", "automobile", "bird", "cat", "deer",
                "dog", "frog", "horse", "ship", "truck"
            ]
            self.id2label = {i: label for i, label in enumerate(labels)}
            self.label2id = {label: i for i, label in enumerate(labels)}
        else:
            self.id2label = id2label
            self.label2id = label2id 
        self.id2label = {int(k): v for k, v in self.id2label.items()}
        self._val_invars()

    def _val_invars(self):
        if self.num_classes <= 0:
            raise ValueError(f"num_classes must be positive, got {self.num_classes}")
        if self.classifier_hidden_dim <= 0:
            raise ValueError(
                f"classifier_hidden_dim must be positive, got {self.classifier_hidden_dim}"
            )
        if len(self.input_size) != 3:
            raise ValueError(f"input_size must be (C, H, W), got {self.input_size}")

        C, H, W = self.input_size
        if C != self.in_channels:
            raise ValueError(f"channel mismatch: input_size[0]={C}, in_channels={self.in_channels}")
        if H <= 0 or W <= 0:
            raise ValueError(f"spatial dimensions must be positive, got H={H}, W={W}")
            
        if len(self.block_channels) == 0:
            raise ValueError("the model needs at least one layer!")
        if any(c <= 0 for c in self.block_channels):
            raise ValueError(f"invalid block_channels: {self.block_channels}")
        if self.use_se and self.se_reduction <= 0:
            raise ValueError(f"se_reduction must be positive, got {self.se_reduction}")
        if len(self.mean) != self.in_channels or len(self.std) != self.in_channels:
            raise ValueError("mean/std length must match in_channels")
            
        if self.num_classes != len(self.id2label):
            raise ValueError(f"num_classes ({self.num_classes}) != label count ({len(self.id2label)})")
        if set(self.label2id.values()) != set(range(self.num_classes)):
            raise ValueError("label2id values must cover [0, ..., num_classes-1]")