Views
No views yet

openai/clip-vit-base-patch32pip install transformers torch safetensors1from transformers import CLIPVisionConfig, CLIPVisionModel, CLIPFeatureExtractor
2import torch
3from torch import nn
4
5class PCamClassifier(nn.Module):
6 def __init__(self, config_dict):
7 super().__init__()
8 self.config = CLIPVisionConfig(**config_dict)
9 self.vision_model = CLIPVisionModel(self.config)
10 self.classifier = nn.Linear(self.config.hidden_size, 2)
11
12 def forward(self, pixel_values):
13 outputs = self.vision_model(pixel_values)
14 return self.classifier(outputs.pooler_output)
15
16# Load model
17config_dict = {
18 "_name_or_path": "openai/clip-vit-base-patch32",
19 "architectures": ["CLIPVisionModel"],
20 "attention_dropout": 0.0,
21 "dropout": 0.0,
22 "hidden_act": "quick_gelu",
23 "hidden_size": 768,
24 "image_size": 224,
25 "initializer_factor": 1.0,
26 "initializer_range": 0.02,
27 "intermediate_size": 3072,
28 "layer_norm_eps": 1e-05,
29 "model_type": "clip_vision_model",
30 "num_attention_heads": 12,
31 "num_channels": 3,
32 "num_hidden_layers": 12,
33 "patch_size": 32,
34 "projection_dim": 512,
35 "torch_dtype": "float32"
36}
37
38# Initialize model
39model = PCamClassifier(config_dict)
40model.load_state_dict(torch.load('best_enhanced_pcam_model.pt'))
41
42
43class PCamDataset(Dataset):
44 def __init__(self, dataset):
45 self.dataset = dataset
46
47 def __len__(self):
48 return len(self.dataset)
49
50 def __getitem__(self, idx):
51 example = self.dataset[idx]
52 image = example["image"].convert("RGB")
53 image_array = np.array(image) / 255.0
54 image_array = image_array.transpose(2, 0, 1).astype(np.float32)
55 return {
56 "pixel_values": image_array,
57 "labels": example["label"]
58 }
59