ResNet / configuration_myresnet.py
JangTaeng's picture
Upload 8 files
b1fe7e3 verified
Raw
History Blame Contribute Delete
2.39 kB
"""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)}")