File size: 4,757 Bytes
29271db | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 | # Model Architecture
**AI vs Real Image Detector β 3-Branch Ensemble**
Input: RGB image resized to **224Γ224**, pixels in `[0, 1]`.
Output: 2 logits β softmax. **Label map: `0 = real`, `1 = ai`.**
```
βββββββββββββββββββββββββββββββ
224Γ224 RGB ββββββββΊβ Branch 1: CLIP (frozen) ββββΊ 768
β βββββββββββββββββββββββββββββββ
β βββββββββββββββββββββββββββββββ
ββββββββββββββββΊβ Branch 2: EfficientNet-B3 ββββΊ 1536
β βββββββββββββββββββββββββββββββ
β βββββββββββββββββββββββββββββββ
ββββββββββββββββΊβ Branch 3: FFT-CNN ββββΊ 512
βββββββββββββββββββββββββββββββ
β concat (2816)
βΌ
βββββββββββββββββββββββββββββββ
β Fusion MLP β
β 2816 β 512 β 128 β 2 β
βββββββββββββββββββββββββββββββ
βΌ
[P(real), P(ai)]
```
---
## Branch 1 β CLIP (semantic view)
- **Backbone:** OpenAI **CLIP ViT-L-14** visual encoder, **fully frozen**
(`requires_grad = False`), run in `eval()` under `torch.no_grad()`.
- **Normalization:** CLIP's own mean/std, applied as registered buffers.
- **Adapter (trainable):** `Linear(768β768) β LayerNorm β GELU β Dropout(0.1) β
Linear(768β768) β LayerNorm`.
- **Output:** 768-d.
- **Rationale:** CLIP encodes high-level semantic structure; the adapter
re-projects those features toward the real-vs-AI decision without disturbing
the pretrained encoder (parameter-efficient, overfitting-resistant).
## Branch 2 β EfficientNet-B3 (texture view)
- **Backbone:** torchvision **EfficientNet-B3**, ImageNet-pretrained,
**fine-tuned** (features + avgpool).
- **Normalization:** ImageNet mean/std (registered buffers).
- **Head:** Dropout(0.4) after global average pooling.
- **Output:** 1536-d.
- **Rationale:** captures local texture and spatial artifacts β skin, edges,
fine detail β where generators leave subtle inconsistencies.
## Branch 3 β FFT-CNN (frequency view)
- **Spectral transform:** convert to grayscale (luma), 2D FFT, `log1p(|Β·|)`
magnitude, `fftshift`, then per-image min-max normalize to `[0, 1]`.
- **CNN:** 4 conv blocks `1β32β64β128β256` (3Γ3, BN, ReLU, MaxPool between the
first three), `AdaptiveAvgPool2d(1)`, flatten.
- **Projection:** `Linear(256β512) β LayerNorm β GELU`.
- **Output:** 512-d.
- **Rationale:** GAN/diffusion upsampling leaves periodic grid patterns in the
frequency domain that are invisible in pixel space. This branch reads them
directly. (Note: this is also why *resizing* an image β which injects
interpolation artifacts β can fool the branch; see limitations.)
## Fusion head
- Concatenate `[768 + 1536 + 512] = 2816`.
- `Linear(2816β512) β BN β GELU β Dropout(0.4) β Linear(512β128) β BN β GELU β
Dropout(0.2) β Linear(128β2)`.
## Parameters
| | Count |
|---|---|
| Total | ~317,877,130 |
| Trainable | ~13,910,922 |
Trainable = CLIP adapter + EfficientNet + FFT branch + fusion head. The CLIP
backbone (the bulk of the params) is frozen.
## Training setup
- **Optimizer:** AdamW with **per-branch learning rates** β CLIP adapter `1e-5`,
EfficientNet `3e-5`, FFT `1e-4`, fusion head `1e-4`; weight decay `1e-4`.
- **Scheduler:** CosineAnnealingLR (`eta_min = 1e-7`).
- **Loss:** class-weighted CrossEntropy. **Grad clipping:** `1.0`.
- **Sampling:** WeightedRandomSampler for class balance.
- **Augmentation:** horizontal flip, color jitter, random grayscale, Gaussian
blur.
## Checkpoint format
`best_model.pt` is a dict:
```python
{"epoch": int, "model_state": state_dict, "val_f1": float, "cfg": {...}}
```
Load via `model.py`'s `load_detector(path, device)`, which rebuilds the network
from the embedded `cfg` and loads `model_state`. `AutoModel`/`pipeline()` do
**not** work β this is a custom `nn.Module`.
|