ResNet / modeling_myresnet.py
JangTaeng's picture
Upload 8 files
b1fe7e3 verified
Raw
History Blame Contribute Delete
9.33 kB
"""MyResNet ๋ชจ๋ธ ๊ตฌํ˜„.
ResNet (Deep Residual Learning for Image Recognition, He et al., 2015)
๋…ผ๋ฌธ์„ PyTorch + ํ—ˆ๊น…ํŽ˜์ด์Šค transformers ํฌ๋งท์œผ๋กœ ๊ตฌํ˜„ํ•œ ํŒŒ์ผ์ž…๋‹ˆ๋‹ค.
"""
from typing import Optional, Union, Tuple
import torch
import torch.nn as nn
from transformers import PreTrainedModel
from transformers.modeling_outputs import ImageClassifierOutput
from configuration_myresnet import MyResNetConfig
# ============================================================
# Basic Block (ResNet-18/34์šฉ) - ๋…ผ๋ฌธ Fig 2
# ============================================================
class BasicBlock(nn.Module):
"""2๊ฐœ์˜ 3x3 conv๋กœ ์ด๋ฃจ์–ด์ง„ ๊ธฐ๋ณธ residual block.
y = ReLU( BN(conv(ReLU(BN(conv(x))))) + shortcut(x) )
"""
expansion = 1 # ์ถœ๋ ฅ ์ฑ„๋„ ๋ฐฐ์ˆ˜
def __init__(self, in_channels: int, out_channels: int, stride: int = 1):
super().__init__()
# ์ฒซ ๋ฒˆ์งธ 3x3 conv (stride๋กœ ๋‹ค์šด์ƒ˜ํ”Œ๋ง ๊ฐ€๋Šฅ)
self.conv1 = nn.Conv2d(
in_channels, out_channels, kernel_size=3, stride=stride, padding=1, bias=False
)
self.bn1 = nn.BatchNorm2d(out_channels)
# ๋‘ ๋ฒˆ์งธ 3x3 conv (stride=1 ๊ณ ์ •)
self.conv2 = nn.Conv2d(
out_channels, out_channels, kernel_size=3, stride=1, padding=1, bias=False
)
self.bn2 = nn.BatchNorm2d(out_channels)
self.relu = nn.ReLU(inplace=True)
# Shortcut: ์ฐจ์›์ด ๋ฐ”๋€” ๋•Œ๋งŒ projection (1x1 conv) ์‚ฌ์šฉ (๋…ผ๋ฌธ Eqn.2, Option B)
if stride != 1 or in_channels != out_channels * self.expansion:
self.shortcut = nn.Sequential(
nn.Conv2d(
in_channels,
out_channels * self.expansion,
kernel_size=1,
stride=stride,
bias=False,
),
nn.BatchNorm2d(out_channels * self.expansion),
)
else:
self.shortcut = nn.Identity()
def forward(self, x: torch.Tensor) -> torch.Tensor:
identity = self.shortcut(x)
out = self.relu(self.bn1(self.conv1(x)))
out = self.bn2(self.conv2(out))
out = out + identity # ๋…ผ๋ฌธ Eqn.1: F(x) + x
return self.relu(out) # addition ํ›„ ReLU (๋…ผ๋ฌธ Sec 3.2)
# ============================================================
# Bottleneck Block (ResNet-50/101/152์šฉ) - ๋…ผ๋ฌธ Fig 5 ์˜ค๋ฅธ์ชฝ
# ============================================================
class BottleneckBlock(nn.Module):
"""1x1 -> 3x3 -> 1x1 ๊ตฌ์กฐ๋กœ ์ฑ„๋„์„ ์ค„์˜€๋‹ค๊ฐ€ ๋‹ค์‹œ ๋Š˜๋ฆฌ๋Š” bottleneck block.
๋งˆ์ง€๋ง‰ 1x1 conv์—์„œ ์ฑ„๋„์„ expansion(=4)๋ฐฐ๋กœ ํ™•์žฅํ•ฉ๋‹ˆ๋‹ค.
"""
expansion = 4
def __init__(self, in_channels: int, out_channels: int, stride: int = 1):
super().__init__()
# 1x1 conv: ์ฑ„๋„ ์ถ•์†Œ
self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=1, bias=False)
self.bn1 = nn.BatchNorm2d(out_channels)
# 3x3 conv: ์‹ค์ œ ์—ฐ์‚ฐ (bottleneck ์ค‘์‹ฌ)
self.conv2 = nn.Conv2d(
out_channels, out_channels, kernel_size=3, stride=stride, padding=1, bias=False
)
self.bn2 = nn.BatchNorm2d(out_channels)
# 1x1 conv: ์ฑ„๋„ ๋ณต์› (4๋ฐฐ๋กœ ํ™•์žฅ)
self.conv3 = nn.Conv2d(
out_channels, out_channels * self.expansion, kernel_size=1, bias=False
)
self.bn3 = nn.BatchNorm2d(out_channels * self.expansion)
self.relu = nn.ReLU(inplace=True)
if stride != 1 or in_channels != out_channels * self.expansion:
self.shortcut = nn.Sequential(
nn.Conv2d(
in_channels,
out_channels * self.expansion,
kernel_size=1,
stride=stride,
bias=False,
),
nn.BatchNorm2d(out_channels * self.expansion),
)
else:
self.shortcut = nn.Identity()
def forward(self, x: torch.Tensor) -> torch.Tensor:
identity = self.shortcut(x)
out = self.relu(self.bn1(self.conv1(x)))
out = self.relu(self.bn2(self.conv2(out)))
out = self.bn3(self.conv3(out))
out = out + identity
return self.relu(out)
# ============================================================
# PreTrainedModel ๋ฒ ์ด์Šค ํด๋ž˜์Šค
# ============================================================
class MyResNetPreTrainedModel(PreTrainedModel):
"""from_pretrained / save_pretrained ๋“ฑ์„ ์ง€์›ํ•˜๋Š” ๋ฒ ์ด์Šค ํด๋ž˜์Šค."""
config_class = MyResNetConfig
base_model_prefix = "myresnet"
main_input_name = "pixel_values"
supports_gradient_checkpointing = False
def _init_weights(self, module):
"""๋…ผ๋ฌธ Sec 3.4: He initialization ์‚ฌ์šฉ."""
if isinstance(module, nn.Conv2d):
nn.init.kaiming_normal_(module.weight, mode="fan_out", nonlinearity="relu")
elif isinstance(module, nn.BatchNorm2d):
nn.init.constant_(module.weight, 1)
nn.init.constant_(module.bias, 0)
elif isinstance(module, nn.Linear):
nn.init.normal_(module.weight, mean=0.0, std=0.01)
if module.bias is not None:
nn.init.constant_(module.bias, 0)
# ============================================================
# ์ด๋ฏธ์ง€ ๋ถ„๋ฅ˜์šฉ ResNet
# ============================================================
class MyResNetForImageClassification(MyResNetPreTrainedModel):
"""์ด๋ฏธ์ง€ ๋ถ„๋ฅ˜์šฉ ResNet ๋ชจ๋ธ.
์ž…๋ ฅ: pixel_values (batch, num_channels, H, W)
์ถœ๋ ฅ: ImageClassifierOutput(loss, logits)
๋…ผ๋ฌธ Table 1์˜ ๊ตฌ์กฐ๋ฅผ ๋”ฐ๋ฆ…๋‹ˆ๋‹ค:
conv1 (7x7, stride=2) -> maxpool -> stage1~4 -> avgpool -> fc
"""
def __init__(self, config: MyResNetConfig):
super().__init__(config)
self.config = config
# ๋ธ”๋ก ์ข…๋ฅ˜ ์„ ํƒ
block = BasicBlock if config.block_type == "basic" else BottleneckBlock
self.in_channels = 64
# ---- Stem (๋…ผ๋ฌธ Table 1 conv1) ----
# 7x7 conv, 64 channels, stride=2 + BN + ReLU + maxpool
self.stem = nn.Sequential(
nn.Conv2d(
config.num_channels, 64, kernel_size=7, stride=2, padding=3, bias=False
),
nn.BatchNorm2d(64),
nn.ReLU(inplace=True),
nn.MaxPool2d(kernel_size=3, stride=2, padding=1),
)
# ---- 4 stages (conv2_x ~ conv5_x) ----
self.stage1 = self._make_stage(
block, config.hidden_sizes[0], config.layers[0], stride=1
)
self.stage2 = self._make_stage(
block, config.hidden_sizes[1], config.layers[1], stride=2
)
self.stage3 = self._make_stage(
block, config.hidden_sizes[2], config.layers[2], stride=2
)
self.stage4 = self._make_stage(
block, config.hidden_sizes[3], config.layers[3], stride=2
)
# ---- Classification head ----
self.avgpool = nn.AdaptiveAvgPool2d(output_size=1)
self.classifier = nn.Linear(
config.hidden_sizes[3] * block.expansion, config.num_labels
)
# ๊ฐ€์ค‘์น˜ ์ดˆ๊ธฐํ™” ์ ์šฉ
self.post_init()
def _make_stage(self, block, out_channels: int, num_blocks: int, stride: int):
"""ํ•˜๋‚˜์˜ stage(๋™์ผ ํ•ด์ƒ๋„์˜ ๋ธ”๋ก๋“ค)๋ฅผ ๋งŒ๋“œ๋Š” ํ—ฌํผ.
์ฒซ ๋ธ”๋ก๋งŒ stride๋กœ ๋‹ค์šด์ƒ˜ํ”Œ๋งํ•˜๊ณ , ๋‚˜๋จธ์ง€๋Š” stride=1.
"""
strides = [stride] + [1] * (num_blocks - 1)
layers = []
for s in strides:
layers.append(block(self.in_channels, out_channels, stride=s))
self.in_channels = out_channels * block.expansion
return nn.Sequential(*layers)
def forward(
self,
pixel_values: torch.Tensor,
labels: Optional[torch.Tensor] = None,
return_dict: Optional[bool] = None,
) -> Union[Tuple, ImageClassifierOutput]:
"""์ˆœ์ „ํŒŒ.
Args:
pixel_values: (batch, num_channels, H, W) ํ˜•ํƒœ์˜ ์ด๋ฏธ์ง€ ํ…์„œ.
labels: (batch,) ํ˜•ํƒœ์˜ ์ •๋‹ต ๋ ˆ์ด๋ธ”. ์ฃผ์–ด์ง€๋ฉด loss๋„ ํ•จ๊ป˜ ๋ฐ˜ํ™˜.
return_dict: True๋ฉด ImageClassifierOutput, False๋ฉด ํŠœํ”Œ ๋ฐ˜ํ™˜.
"""
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
# Shape ํ๋ฆ„: (B, 3, 224, 224)
x = self.stem(pixel_values) # -> (B, 64, 56, 56)
x = self.stage1(x) # -> (B, 64*e, 56, 56)
x = self.stage2(x) # -> (B, 128*e, 28, 28)
x = self.stage3(x) # -> (B, 256*e, 14, 14)
x = self.stage4(x) # -> (B, 512*e, 7, 7)
x = self.avgpool(x) # -> (B, 512*e, 1, 1)
x = torch.flatten(x, 1) # -> (B, 512*e)
logits = self.classifier(x) # -> (B, num_labels)
# Loss ๊ณ„์‚ฐ (Trainer ํ˜ธํ™˜)
loss = None
if labels is not None:
loss_fn = nn.CrossEntropyLoss()
loss = loss_fn(logits, labels)
if not return_dict:
output = (logits,)
return ((loss,) + output) if loss is not None else output
return ImageClassifierOutput(loss=loss, logits=logits)