Views
No views yet
FaceForge Generator (252.5M parameters)
│
├── ViT Encoders (172M params)
│ ├── Source Encoder: ViT-B/16 (86M)
│ │ └── 12 layers, 768-dim, 12 heads
│ └── Target Encoder: ViT-B/16 (86M)
│ └── 12 layers, 768-dim, 12 heads
│
├── Cross-Attention Module (14M params)
│ ├── 2 layers, 8 heads
│ ├── FFN: 768 → 3072 → 768
│ └── Dropout: 0.1
│
├── Transformer Decoder (58M params)
│ ├── 256 learnable queries (16×16)
│ ├── 6 decoder layers, 8 heads
│ └── 2D positional embeddings
│
└── CNN Upsampler (9M params)
├── TransposeConv: 768→512→256→128→64
├── 4 upsampling stages (16×16 → 224×224)
└── Conv: 64→32→3 + Tanh| Epoch | Train Loss | Val Loss | Time (min) |
|---|---|---|---|
| 1 | 0.2873 | 0.2804 | 227.5 |
| 2 | 0.2432 | 0.2304 | 231.2 |
| 3 | 0.2143 | 0.2043 | 228.8 |
pip install torch torchvision timm pillow numpy1import torch
2import torch.nn as nn
3import timm
4from torchvision import transforms
5
6class FaceForgeGenerator(nn.Module):
7 def __init__(self):
8 super().__init__()
9 # Source and Target ViT Encoders
10 self.source_encoder = timm.create_model('vit_base_patch16_224', pretrained=True, num_classes=0)
11 self.target_encoder = timm.create_model('vit_base_patch16_224', pretrained=True, num_classes=0)
12
13 # Cross-attention (implement your architecture)
14 # Transformer decoder
15 # CNN upsampler
16 # ... (see full architecture in paper)
17
18 def forward(self, source_face, target_face):
19 # Encode both faces
20 source_features = self.source_encoder.forward_features(source_face)
21 target_features = self.target_encoder.forward_features(target_face)
22
23 # Cross-attention fusion
24 fused_features = self.cross_attention(source_features, target_features)
25
26 # Decode to spatial map
27 spatial_features = self.transformer_decoder(fused_features)
28
29 # Upsample to 224×224
30 generated_face = self.cnn_upsampler(spatial_features)
31
32 return generated_face
33
34# Load checkpoint
35model = FaceForgeGenerator()
36checkpoint = torch.load('generator_best.pth', map_location='cpu')
37model.load_state_dict(checkpoint['model_state_dict'])
38model.eval()
39
40# Preprocessing
41transform = transforms.Compose([
42 transforms.Resize((224, 224)),
43 transforms.ToTensor(),
44 transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])
45])
46
47# Generate face swap
48def generate_face_swap(source_path, target_path):
49 source = transform(Image.open(source_path).convert('RGB')).unsqueeze(0)
50 target = transform(Image.open(target_path).convert('RGB')).unsqueeze(0)
51
52 with torch.no_grad():
53 generated = model(source, target)
54
55 # Denormalize and convert to PIL
56 generated = (generated[0] * 0.5 + 0.5).clamp(0, 1)
57 generated = transforms.ToPILImage()(generated)
58
59 return generated
60
61# Example
62result = generate_face_swap("source.jpg", "target.jpg")
63result.save("generated.jpg")1optimizer: AdamW
2learning_rate: 1e-4
3betas: [0.9, 0.999]
4weight_decay: 1e-4
5batch_size: 16
6epochs: 3 (baseline)
7loss_function: L1 (Mean Absolute Error)
8lr_schedule: Cosine Annealing (1e-4 → 1e-6)1@techreport{nasir2026faceforge,
2 title={FaceForge: A Deep Learning Framework for Facial Manipulation Generation and Detection},
3 author={Nasir, Huzaifa},
4 institution={National University of Computer and Emerging Sciences},
5 year={2026},
6 doi={10.5281/zenodo.18530439}
7}