# from torchvision import datasets, transforms, models
# from torch.utils.data import DataLoader
# import torch, torch.nn as nn, torch.optim as optim

# DATA_DIR = "data_soil_visibility"
# MODEL_PATH = "soil_visibility_classifier.pt"

# transform = transforms.Compose([
#     transforms.Resize((224, 224)),
#     transforms.RandomHorizontalFlip(),
#     transforms.RandomRotation(10),
#     transforms.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.2),
#     transforms.ToTensor(),
#     transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
# ])

# dataset = datasets.ImageFolder(DATA_DIR, transform=transform)
# dataloader = DataLoader(dataset, batch_size=16, shuffle=True)

# model = models.resnet18(pretrained=True)
# model.fc = nn.Linear(model.fc.in_features, 1)
# device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# model = model.to(device)

# criterion = nn.BCEWithLogitsLoss()
# optimizer = optim.Adam(model.parameters(), lr=0.001)

# for epoch in range(10):
#     model.train()
#     for inputs, labels in dataloader:
#         inputs, labels = inputs.to(device), labels.float().unsqueeze(1).to(device)
#         optimizer.zero_grad()
#         outputs = model(inputs)
#         loss = criterion(outputs, labels)
#         loss.backward()
#         optimizer.step()

# torch.save(model.state_dict(), MODEL_PATH)
# print(f"✅ Saved weights to {MODEL_PATH}")


from torchvision import datasets, transforms
from torch.utils.data import DataLoader
import torch, torch.nn as nn, torch.optim as optim
import timm  # EfficientNet comes from timm

DATA_DIR = "data_soil_visibility"
MODEL_PATH = "soil_visibility_classifier.pt"

# Transforms
transform = transforms.Compose([
    transforms.Resize((224, 224)),
    transforms.RandomHorizontalFlip(),
    transforms.RandomRotation(10),
    transforms.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.2),
    transforms.ToTensor(),
    transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
])

# Dataset + loader
dataset = datasets.ImageFolder(DATA_DIR, transform=transform)
dataloader = DataLoader(dataset, batch_size=16, shuffle=True)

# EfficientNet-B0 backbone
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = timm.create_model("efficientnet_b0", pretrained=True, num_classes=1)
model = model.to(device)

# Loss + optimizer
criterion = nn.BCEWithLogitsLoss()
optimizer = optim.Adam(model.parameters(), lr=0.001)

# Training loop
for epoch in range(10):
    model.train()
    for inputs, labels in dataloader:
        inputs, labels = inputs.to(device), labels.float().unsqueeze(1).to(device)
        optimizer.zero_grad()
        outputs = model(inputs)
        loss = criterion(outputs, labels)
        loss.backward()
        optimizer.step()

# Save weights only
torch.save(model.state_dict(), MODEL_PATH)
print(f"✅ Saved EfficientNet-B0 weights to {MODEL_PATH}")
