| import torch
|
| import torch.nn as nn
|
| from torch.utils.data import DataLoader, ConcatDataset
|
| from dataset import COCOSegmentationDataset
|
| from model import ResNetSegmentation
|
|
|
| import os
|
|
|
|
|
|
|
| DATA_ROOT = "data"
|
| BATCH_SIZE = 8
|
| EPOCHS = 5
|
| LR = 1e-4
|
|
|
| DEVICE = "cpu"
|
| torch.set_num_threads(8)
|
|
|
| print("Device:", DEVICE)
|
| print("⚠️ RTX 5070 (sm_120) requires PyTorch nightly build or future release")
|
|
|
|
|
|
|
| crack_root = os.path.join(DATA_ROOT, "cracks.v1-cracks-f.coco")
|
| drywall_root = os.path.join(DATA_ROOT, "Drywall-Join-Detect.v2i.coco")
|
|
|
| train_full_dataset = ConcatDataset([
|
| COCOSegmentationDataset(crack_root, "train"),
|
| COCOSegmentationDataset(drywall_root, "train"),
|
| ])
|
|
|
|
|
|
|
|
|
|
|
| train_dataset = train_full_dataset
|
|
|
| val_dataset = ConcatDataset([
|
| COCOSegmentationDataset(crack_root, "valid"),
|
| COCOSegmentationDataset(drywall_root, "valid"),
|
| ])
|
|
|
| train_loader = DataLoader(
|
| train_dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=2
|
| )
|
| val_loader = DataLoader(
|
| val_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=2
|
| )
|
|
|
| print("Train samples (Limited):", len(train_dataset))
|
| print("Val samples:", len(val_dataset))
|
|
|
|
|
|
|
| model = ResNetSegmentation().to(DEVICE)
|
|
|
| criterion = nn.BCEWithLogitsLoss()
|
| optimizer = torch.optim.AdamW(model.parameters(), lr=LR)
|
|
|
|
|
|
|
| for epoch in range(EPOCHS):
|
| model.train()
|
| train_loss = 0.0
|
|
|
| print(f"Epoch {epoch+1}/{EPOCHS}...")
|
| for step, (images, masks) in enumerate(train_loader):
|
| images = images.to(DEVICE)
|
| masks = masks.to(DEVICE)
|
|
|
| preds = model(images)
|
| loss = criterion(preds, masks)
|
|
|
| optimizer.zero_grad()
|
| loss.backward()
|
| optimizer.step()
|
|
|
| train_loss += loss.item()
|
| if step % 10 == 0:
|
| print(f" Step {step}/{len(train_loader)} Loss: {loss.item():.4f}")
|
|
|
| avg_loss = train_loss / len(train_loader)
|
| print(f"Epoch {epoch+1} Train Loss: {avg_loss:.4f}")
|
|
|
|
|
| MODEL_PATH = "best_model.pth"
|
| torch.save(model.state_dict(), MODEL_PATH)
|
| print(f"✅ Model saved to: {MODEL_PATH}")
|
|
|
| print(f"✅ Training completed!")
|
|
|