| """Food-101 ๋ฐ์ดํฐ์
์ผ๋ก MyResNet์ ํ์ตํ๋ ์คํฌ๋ฆฝํธ. |
| |
| torchvision์ ImageNet pretrained ResNet-18 ๊ฐ์ค์น๋ฅผ ๊ฐ์ ธ์์ |
| MyResNet์ ๋ก๋ํ ํ Food-101์ fine-tuningํฉ๋๋ค. |
| |
| ์ฌ์ฉ๋ฒ: |
| python train_food.py |
| |
| ์๊ตฌ์ฌํญ: |
| pip install torch torchvision transformers datasets accelerate |
| |
| ํ์ต ์๊ฐ (GPU 1์ฅ ๊ธฐ์ค): |
| - Food-101 full (101์ข
): ์ฝ 2~4์๊ฐ (10 epoch) |
| - ๋น ๋ฅธ ํ
์คํธ: epochs=3์ผ๋ก ์ค์ด๋ฉด 1์๊ฐ ๋ด |
| """ |
| import numpy as np |
| import torch |
| import torch.nn as nn |
| import torchvision.models as tv_models |
| from datasets import load_dataset |
| from torchvision.transforms import ( |
| Compose, |
| Normalize, |
| RandomCrop, |
| RandomHorizontalFlip, |
| Resize, |
| ToTensor, |
| ) |
| from transformers import DefaultDataCollator, Trainer, TrainingArguments |
|
|
| from configuration_myresnet import MyResNetConfig |
| from modeling_myresnet import MyResNetForImageClassification |
|
|
|
|
| |
| |
| |
| print("Food-101 ๋ฐ์ดํฐ์
๋ก๋ฉ ์ค...") |
| dataset = load_dataset("food101") |
|
|
| |
| food_classes = dataset["train"].features["label"].names |
| NUM_CLASSES = len(food_classes) |
| print(f"์ด ํด๋์ค ์: {NUM_CLASSES}") |
| print(f"ํ์ต ์ด๋ฏธ์ง: {len(dataset['train']):,}์ฅ") |
| print(f"๊ฒ์ฆ ์ด๋ฏธ์ง: {len(dataset['validation']):,}์ฅ") |
|
|
|
|
| |
| |
| |
| config = MyResNetConfig( |
| num_channels=3, |
| num_labels=NUM_CLASSES, |
| block_type="basic", |
| layers=[2, 2, 2, 2], |
| hidden_sizes=[64, 128, 256, 512], |
| image_size=224, |
| id2label={i: name for i, name in enumerate(food_classes)}, |
| label2id={name: i for i, name in enumerate(food_classes)}, |
| ) |
| model = MyResNetForImageClassification(config) |
|
|
|
|
| def load_pretrained_resnet18(model): |
| """torchvision์ ImageNet pretrained ResNet-18 ๊ฐ์ค์น๋ฅผ MyResNet์ ๋ณต์ฌ. |
| |
| ๋ ๋ชจ๋ธ์ ๋ ์ด์ด ์ด๋ฆ์ด ๋ค๋ฅด๋ฏ๋ก ๋งคํํด์ค๋๋ค. |
| ๋ง์ง๋ง FC ๋ ์ด์ด(classifier)๋ ํด๋์ค ์๊ฐ ๋ค๋ฅด๋ฏ๋ก ์คํต. |
| """ |
| print("ImageNet pretrained ResNet-18 ๊ฐ์ค์น ๋ก๋ฉ...") |
| tv_resnet = tv_models.resnet18(weights=tv_models.ResNet18_Weights.IMAGENET1K_V1) |
| tv_state = tv_resnet.state_dict() |
|
|
| |
| |
| |
| mapping = { |
| "conv1.": "stem.0.", |
| "bn1.": "stem.1.", |
| "layer1.": "stage1.", |
| "layer2.": "stage2.", |
| "layer3.": "stage3.", |
| "layer4.": "stage4.", |
| } |
|
|
| |
| new_state = {} |
| for k, v in tv_state.items(): |
| if k.startswith("fc."): |
| continue |
|
|
| new_k = k |
| for old, new in mapping.items(): |
| if new_k.startswith(old): |
| new_k = new_k.replace(old, new, 1) |
| break |
|
|
| |
| new_k = new_k.replace(".downsample.", ".shortcut.") |
| new_state[new_k] = v |
|
|
| |
| missing, unexpected = model.load_state_dict(new_state, strict=False) |
| print(f"๋ก๋ ์ฑ๊ณต. ๋๋ฝ๋ ํค {len(missing)}๊ฐ (classifier ๋ฑ), " |
| f"์์์น ๋ชปํ ํค {len(unexpected)}๊ฐ") |
| return model |
|
|
|
|
| model = load_pretrained_resnet18(model) |
| print(f"๋ชจ๋ธ ํ๋ผ๋ฏธํฐ ์: {sum(p.numel() for p in model.parameters()):,}") |
|
|
|
|
| |
| |
| |
| IMAGENET_MEAN = [0.485, 0.456, 0.406] |
| IMAGENET_STD = [0.229, 0.224, 0.225] |
|
|
| train_transform = Compose([ |
| Resize(256), |
| RandomCrop(224), |
| RandomHorizontalFlip(), |
| ToTensor(), |
| Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD), |
| ]) |
| eval_transform = Compose([ |
| Resize((224, 224)), |
| ToTensor(), |
| Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD), |
| ]) |
|
|
|
|
| def preprocess_train(batch): |
| batch["pixel_values"] = [ |
| train_transform(img.convert("RGB")) for img in batch["image"] |
| ] |
| batch["labels"] = batch["label"] |
| return batch |
|
|
|
|
| def preprocess_eval(batch): |
| batch["pixel_values"] = [ |
| eval_transform(img.convert("RGB")) for img in batch["image"] |
| ] |
| batch["labels"] = batch["label"] |
| return batch |
|
|
|
|
| train_ds = dataset["train"].with_transform(preprocess_train) |
| eval_ds = dataset["validation"].with_transform(preprocess_eval) |
|
|
|
|
| |
| |
| |
| def compute_metrics(eval_pred): |
| logits, labels = eval_pred |
| |
| top1_preds = np.argmax(logits, axis=-1) |
| top1_acc = (top1_preds == labels).mean() |
|
|
| |
| top5_preds = np.argsort(-logits, axis=-1)[:, :5] |
| top5_correct = np.any(top5_preds == labels[:, None], axis=-1) |
| top5_acc = top5_correct.mean() |
|
|
| return { |
| "accuracy": float(top1_acc), |
| "top5_accuracy": float(top5_acc), |
| } |
|
|
|
|
| |
| |
| |
| training_args = TrainingArguments( |
| output_dir="./my-resnet18-food101", |
| num_train_epochs=10, |
| per_device_train_batch_size=64, |
| per_device_eval_batch_size=64, |
| learning_rate=1e-3, |
| weight_decay=1e-4, |
| lr_scheduler_type="cosine", |
| warmup_ratio=0.05, |
| eval_strategy="epoch", |
| save_strategy="epoch", |
| save_total_limit=2, |
| logging_steps=50, |
| load_best_model_at_end=True, |
| metric_for_best_model="accuracy", |
| greater_is_better=True, |
| fp16=torch.cuda.is_available(), |
| remove_unused_columns=False, |
| dataloader_num_workers=4, |
| report_to="none", |
| push_to_hub=False, |
| |
| ) |
|
|
| trainer = Trainer( |
| model=model, |
| args=training_args, |
| train_dataset=train_ds, |
| eval_dataset=eval_ds, |
| data_collator=DefaultDataCollator(), |
| compute_metrics=compute_metrics, |
| ) |
|
|
|
|
| |
| |
| |
| if __name__ == "__main__": |
| print("\n=== Food-101 ํ์ต ์์ ===") |
| trainer.train() |
|
|
| metrics = trainer.evaluate() |
| print(f"\n์ต์ข
Top-1 ์ ํ๋: {metrics['eval_accuracy']:.4f}") |
| print(f"์ต์ข
Top-5 ์ ํ๋: {metrics['eval_top5_accuracy']:.4f}") |
|
|
| trainer.save_model("./my-resnet18-food101") |
| print("๋ชจ๋ธ ์ ์ฅ ์๋ฃ: ./my-resnet18-food101") |
|
|
| |
| |
|
|