| """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 |
|
|
| |
| REPO_ID = "smilelab/visual-valence-model" |
|
|
| CHECKPOINT_FILENAMES = { |
| "stage0": "vca_ckvideo_batch128_lr2e-5_epoch20.pth", |
| "stage1": "vca_IAPS_batch10_lr2e-4_epoch23.pth", |
| "pre": "base_model_vca_IAPS_quadrant.pth", |
| "post_epoch1": "base_model_conditioned_orientation_epoch1.pth", |
| "post": "base_model_conditioned_orientation_epoch100.pth", |
| } |
|
|
| |
| |
| 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() |
|
|