File size: 6,781 Bytes
ddedd58
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
"""Example inference script for the Visual-Valence Model (VCA).

Requirements:
    - Clone https://github.com/lab-smile/FearConditioningAI and run this script from
      inside that repo (or add it to PYTHONPATH), so `models.VGG_Model` / `utils` are
      importable. Install its dependencies (environment-<platform>.yml or requirements.txt).
    - pip install huggingface_hub pillow torch torchvision

Usage:
    python inference_example.py --image path/to/scene.jpg --checkpoint post
    python inference_example.py --gabor path/to/gabor_patch.png --checkpoint post
    python inference_example.py --image path/to/scene.jpg --checkpoint stage1  # full-frame checkpoint

Checkpoints (see CHECKPOINT_FILENAMES): stage0, stage1 (full-frame, Stages 0-1),
pre, post_epoch1, post (quadrant-cropped, Stages 2-3).
"""

import argparse

import torch
from huggingface_hub import hf_hub_download
from PIL import Image
from torchvision import transforms

from models.VGG_Model import Visual_Cortex_Amygdala

# Placeholder — update to the actual Hugging Face repo id once uploaded.
REPO_ID = "smilelab/visual-valence-model"

CHECKPOINT_FILENAMES = {
    "stage0": "vca_ckvideo_batch128_lr2e-5_epoch20.pth",  # trained from scratch on Videoframe
    "stage1": "vca_IAPS_batch10_lr2e-4_epoch23.pth",  # fine-tuned on full-size IAPS
    "pre": "base_model_vca_IAPS_quadrant.pth",  # Stage 2: quadrant fine-tune, before conditioning
    "post_epoch1": "base_model_conditioned_orientation_epoch1.pth",  # Stage 3, epoch 1 (early snapshot)
    "post": "base_model_conditioned_orientation_epoch100.pth",  # Stage 3, epoch 100 (final, after conditioning)
}

# Stage 0/1 checkpoints were trained on full-frame images; Stage 2/3 checkpoints expect the
# quadrant-cropped layout (see preprocess_natural_image vs. preprocess_cs_only below).
FULL_FRAME_CHECKPOINTS = {"stage0", "stage1"}

IMAGE_SIZE = 224
NORMALIZE_MEAN = [0.485, 0.456, 0.406]
NORMALIZE_STD = [0.229, 0.224, 0.225]


def load_model(checkpoint: str = "post", device: str = "cpu") -> torch.nn.Module:
    """Download a checkpoint from the Hub and load it into a Visual_Cortex_Amygdala model."""
    filename = CHECKPOINT_FILENAMES[checkpoint]
    ckpt_path = hf_hub_download(repo_id=REPO_ID, filename=filename)

    model = Visual_Cortex_Amygdala()
    checkpoint_dict = torch.load(ckpt_path, map_location=device, weights_only=False)
    model.load_state_dict(checkpoint_dict["state_dict"], strict=False)
    model.to(device)
    model.eval()
    return model


def preprocess_natural_image(image: Image.Image) -> torch.Tensor:
    """Full-frame preprocessing for the Stage 0/1 checkpoints (no quadrant cropping)."""
    transform = transforms.Compose([
        transforms.Resize((IMAGE_SIZE, IMAGE_SIZE)),
        transforms.ToTensor(),
        transforms.Normalize(mean=NORMALIZE_MEAN, std=NORMALIZE_STD),
    ])
    return transform(image.convert("RGB")).unsqueeze(0)


def _place_in_quadrant(patch: Image.Image, quadrant: int, canvas: Image.Image = None) -> Image.Image:
    """Resize `patch` to quarter size and paste it into one quadrant of `canvas` (new blank one if None).

    quadrant: 1 = top-right, 2 = top-left, 3 = bottom-left, 4 = bottom-right.
    """
    if canvas is None:
        canvas = Image.new("RGB", (IMAGE_SIZE, IMAGE_SIZE))
    patch = patch.convert("RGB").resize((IMAGE_SIZE // 2, IMAGE_SIZE // 2))
    positions = {
        1: (IMAGE_SIZE // 2, 0),
        2: (0, 0),
        3: (0, IMAGE_SIZE // 2),
        4: (IMAGE_SIZE // 2, IMAGE_SIZE // 2),
    }
    canvas.paste(patch, positions[quadrant])
    return canvas


def preprocess_quadrant(image: Image.Image, quadrant: int = 4) -> torch.Tensor:
    """Place a natural (US) scene alone into one quadrant of an otherwise-blank canvas.

    Used for Stage 2/3 checkpoints, which were trained on quadrant-cropped US images
    (see utils.Quadrant_Processing). Default quadrant 4 (bottom-right) matches training.
    """
    canvas = _place_in_quadrant(image, quadrant)
    transform = transforms.Compose([
        transforms.ToTensor(),
        transforms.Normalize(mean=NORMALIZE_MEAN, std=NORMALIZE_STD),
    ])
    return transform(canvas).unsqueeze(0)


def preprocess_cs_only(gabor_patch: Image.Image, quadrant: int = 2) -> torch.Tensor:
    """Place a CS Gabor patch alone into one quadrant of an otherwise-blank canvas.

    Mirrors the conditioning-stage input layout (see utils.Quadrant_Processing_Conditioning /
    test_gaborpatches.py): the US quadrant is left blank so the response reflects only what
    the model has learned to associate with the CS. Default quadrant 2 (top-left) matches
    the CS placement used during Stage 3 training; only meaningful for Stage 2/3 checkpoints.
    """
    return preprocess_quadrant(gabor_patch, quadrant)


@torch.no_grad()
def predict_valence(model: torch.nn.Module, input_tensor: torch.Tensor, device: str = "cpu") -> float:
    """Run the model and rescale its sigmoid output from [0, 1] to the [1, 9] IAPS valence scale."""
    output = model(input_tensor.to(device))
    valence = 1 + output.item() * (9 - 1)
    return valence


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--image", type=str, default=None, help="Path to a natural scene image (US-style input).")
    parser.add_argument("--gabor", type=str, default=None, help="Path to a Gabor patch image (CS-only input).")
    parser.add_argument("--checkpoint", choices=list(CHECKPOINT_FILENAMES), default="post",
                        help="Which checkpoint to load (see CHECKPOINT_FILENAMES). Default: 'post' "
                             "(final, post-conditioning model).")
    parser.add_argument("--device", type=str, default="cpu")
    args = parser.parse_args()

    if not args.image and not args.gabor:
        parser.error("Provide --image or --gabor.")
    if args.gabor and args.checkpoint in FULL_FRAME_CHECKPOINTS:
        parser.error(f"--gabor (CS-only input) isn't meaningful for checkpoint '{args.checkpoint}', "
                     f"which was trained on full-frame images without a CS. Use --image instead, or "
                     f"pick a Stage 2/3 checkpoint (pre, post_epoch1, post).")

    model = load_model(checkpoint=args.checkpoint, device=args.device)

    if args.image and args.checkpoint in FULL_FRAME_CHECKPOINTS:
        input_tensor = preprocess_natural_image(Image.open(args.image))
    elif args.image:
        input_tensor = preprocess_quadrant(Image.open(args.image))
    else:
        input_tensor = preprocess_cs_only(Image.open(args.gabor))

    valence = predict_valence(model, input_tensor, device=args.device)
    print(f"Predicted valence (1=extreme displeasure, 9=extreme pleasure): {valence:.2f}")


if __name__ == "__main__":
    main()