| |
| |
| |
| !pip install datasets torch torchvision pillow huggingface_hub |
|
|
| import os |
| import shutil |
| import torch |
| import torch.nn as nn |
| import torch.optim as optim |
| from torch.utils.data import Dataset, DataLoader |
| from torchvision import transforms, models |
| from datasets import load_dataset |
| from PIL import Image |
| from google.colab import files |
|
|
| |
| |
| |
| DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu") |
| MODEL_INPUT_SIZE = 224 |
| SKIN_SIZE = 128 |
| MODEL_PATH = "vit_minecraft_core.pth" |
|
|
| |
| |
| |
| class ViTSkinGenerator(nn.Module): |
| def __init__(self): |
| super().__init__() |
| |
| vit = models.vit_b_16(pretrained=True) |
| self.transformer_core = vit |
| num_features = vit.heads.head.in_features |
| |
| |
| self.transformer_core.heads = nn.Sequential( |
| nn.Linear(num_features, 2048), |
| nn.GELU(), |
| nn.Dropout(0.1), |
| nn.Linear(2048, 4096), |
| nn.GELU(), |
| nn.Linear(4096, 3 * SKIN_SIZE * SKIN_SIZE), |
| nn.Sigmoid() |
| ) |
| |
| def forward(self, x): |
| out = self.transformer_core(x) |
| out = out.view(-1, 3, SKIN_SIZE, SKIN_SIZE) |
| return out |
|
|
| class MinecraftTransformerDataset(Dataset): |
| def __init__(self, img_dir, skin_dir): |
| self.img_dir = img_dir |
| self.skin_dir = skin_dir |
| self.filenames = sorted(os.listdir(img_dir)) |
| self.transform_img = transforms.Compose([ |
| transforms.ToTensor(), |
| transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)) |
| ]) |
| self.transform_skin = transforms.Compose([transforms.ToTensor()]) |
| |
| def __len__(self): return len(self.filenames) |
| def __getitem__(self, idx): |
| f = self.filenames[idx] |
| return self.transform_img(Image.open(os.path.join(self.img_dir, f)).convert("RGB")), \ |
| self.transform_skin(Image.open(os.path.join(self.skin_dir, f)).convert("RGB")) |
|
|
| |
| |
| |
| def train_pipeline(epochs=35, batch_size=16, lr=3e-4, num_samples=437): |
| """Télécharge le dataset, initialise le modèle (ou charge l'existant) et lance l'entraînement.""" |
| print(f"🚀 Initialisation de l'entraînement sur : {DEVICE}") |
| |
| |
| if os.path.exists("dataset"): |
| shutil.rmtree("dataset") |
| os.makedirs("dataset/images", exist_ok=True) |
| os.makedirs("dataset/skins", exist_ok=True) |
|
|
| print("📥 Extraction du dataset Hugging Face...") |
| hf_dataset = load_dataset("zhoudoe23/minecraft-skins-1.1m-prerendered", split="train", streaming=True) |
| count = 0 |
| for item in hf_dataset: |
| if count >= num_samples: break |
| try: |
| render_img = item["rendering_3d"].convert("RGB").resize((MODEL_INPUT_SIZE, MODEL_INPUT_SIZE)) |
| skin_img = item["raw_skin"].convert("RGB").resize((SKIN_SIZE, SKIN_SIZE), Image.NEAREST) |
| render_img.save(f"dataset/images/skin_{count}.png") |
| skin_img.save(f"dataset/skins/skin_{count}.png") |
| count += 1 |
| except: continue |
|
|
| dataset = MinecraftTransformerDataset("dataset/images", "dataset/skins") |
| dataloader = DataLoader(dataset, batch_size=batch_size, shuffle=True) |
|
|
| model = ViTSkinGenerator().to(DEVICE) |
| optimizer = optim.AdamW(model.parameters(), lr=lr, weight_decay=1e-4) |
| criterion = nn.SmoothL1Loss() |
|
|
| if os.path.exists(MODEL_PATH): |
| print(f"💾 Modèle existant trouvé. Reprise de l'apprentissage depuis {MODEL_PATH}...") |
| model.load_state_dict(torch.load(MODEL_PATH, map_location=DEVICE)) |
|
|
| print("\n🏋️ Alignement du Vision Transformer...") |
| for epoch in range(epochs): |
| model.train() |
| epoch_loss = 0 |
| for imgs, skins in dataloader: |
| imgs, skins = imgs.to(DEVICE), skins.to(DEVICE) |
| |
| outputs = model(imgs) |
| loss = criterion(outputs, skins) |
| |
| optimizer.zero_grad() |
| loss.backward() |
| optimizer.step() |
| epoch_loss += loss.item() |
| |
| print(f"Époque [{epoch+1}/{epochs}] | Indice de précision (Loss): {epoch_loss/len(dataloader):.6f}") |
|
|
| torch.save(model.state_dict(), MODEL_PATH) |
| print(f"✅ Modèle sauvegardé avec succès dans : {MODEL_PATH}") |
|
|
| |
| |
| |
| def use_pipeline(image_path, output_path="skin_vit_ultimate.png", download=True): |
| """Prend n'importe quelle image d'entrée, charge les poids du modèle et génère le skin.""" |
| if not os.path.exists(MODEL_PATH): |
| print(f"❌ Erreur : Le fichier de poids '{MODEL_PATH}' n'existe pas. Lancez d'abord l'entraînement.") |
| return |
|
|
| |
| model = ViTSkinGenerator().to(DEVICE) |
| model.load_state_dict(torch.load(MODEL_PATH, map_location=DEVICE)) |
| model.eval() |
|
|
| if not os.path.exists(image_path): |
| print(f"⚠️ Image d'entrée '{image_path}' introuvable.") |
| return |
|
|
| |
| img = Image.open(image_path).convert("RGB").resize((MODEL_INPUT_SIZE, MODEL_INPUT_SIZE), Image.Resampling.LANCZOS) |
| transform = transforms.Compose([ |
| transforms.ToTensor(), |
| transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)) |
| ]) |
| tensor_in = transform(img).unsqueeze(0).to(DEVICE) |
|
|
| with torch.no_grad(): |
| tensor_out = model(tensor_in).squeeze(0).cpu() |
|
|
| |
| tensor_out = torch.clamp(tensor_out, 0, 1) |
| skin_pil = transforms.ToPILImage()(tensor_out) |
| skin_final = skin_pil.resize((SKIN_SIZE, SKIN_SIZE), Image.NEAREST) |
| skin_final.save(output_path) |
| |
| print(f"✨ Skin extrait avec succès dans : {output_path}") |
| if download: |
| files.download(output_path) |
|
|
| |
| |
| |
|
|
| |
| train_pipeline(epochs=25, batch_size=16, num_samples=437) |
|
|
| |
| |
|
|