| """MyResNet ๋ชจ๋ธ ์ค์ ํด๋์ค. |
| |
| ResNet (He et al., 2015) ๋
ผ๋ฌธ์ ๋ฐํ์ผ๋ก ๊ตฌํํ ์ปค์คํ
๋ชจ๋ธ์ ์ค์ ์
๋๋ค. |
| ํ๊น
ํ์ด์ค PretrainedConfig๋ฅผ ์์๋ฐ์ save_pretrained / from_pretrained ํธํ๋ฉ๋๋ค. |
| """ |
| from typing import List, Optional |
|
|
| from transformers import PretrainedConfig |
|
|
|
|
| class MyResNetConfig(PretrainedConfig): |
| """ |
| MyResNet ๋ชจ๋ธ์ ํ์ดํผํ๋ผ๋ฏธํฐ๋ฅผ ์ ์ฅํ๋ Config ํด๋์ค. |
| |
| Args: |
| num_channels (int): ์
๋ ฅ ์ด๋ฏธ์ง ์ฑ๋ ์ (RGB=3, grayscale=1). |
| num_labels (int): ๋ถ๋ฅํ ํด๋์ค ๊ฐ์. |
| block_type (str): 'basic' (ResNet-18/34) ๋๋ 'bottleneck' (ResNet-50/101/152). |
| layers (List[int]): ๊ฐ stage(conv2_x ~ conv5_x)์ ๋ค์ด๊ฐ ๋ธ๋ก ๊ฐ์. |
| - ResNet-18: [2, 2, 2, 2] |
| - ResNet-34: [3, 4, 6, 3] |
| - ResNet-50: [3, 4, 6, 3] (bottleneck) |
| - ResNet-101: [3, 4, 23, 3] (bottleneck) |
| - ResNet-152: [3, 8, 36, 3] (bottleneck) |
| hidden_sizes (List[int]): ๊ฐ stage์ ๊ธฐ๋ณธ ์ฑ๋ ์. ๊ธฐ๋ณธ๊ฐ [64, 128, 256, 512]. |
| image_size (int): ์
๋ ฅ ์ด๋ฏธ์ง ํฌ๊ธฐ (์ ์ฌ๊ฐํ ๊ธฐ์ค). |
| |
| Example: |
| >>> from configuration_myresnet import MyResNetConfig |
| >>> config = MyResNetConfig(num_labels=10, layers=[2, 2, 2, 2]) |
| """ |
|
|
| model_type = "myresnet" |
|
|
| def __init__( |
| self, |
| num_channels: int = 3, |
| num_labels: int = 1000, |
| block_type: str = "basic", |
| layers: Optional[List[int]] = None, |
| hidden_sizes: Optional[List[int]] = None, |
| image_size: int = 224, |
| **kwargs, |
| ): |
| super().__init__(num_labels=num_labels, **kwargs) |
| self.num_channels = num_channels |
| self.block_type = block_type |
| self.layers = layers if layers is not None else [3, 4, 6, 3] |
| self.hidden_sizes = hidden_sizes if hidden_sizes is not None else [64, 128, 256, 512] |
| self.image_size = image_size |
|
|
| |
| if block_type not in ("basic", "bottleneck"): |
| raise ValueError(f"block_type must be 'basic' or 'bottleneck', got {block_type}") |
| if len(self.layers) != 4: |
| raise ValueError(f"layers must have length 4, got {len(self.layers)}") |
| if len(self.hidden_sizes) != 4: |
| raise ValueError(f"hidden_sizes must have length 4, got {len(self.hidden_sizes)}") |
|
|