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()
|