fmegahed/clip
1
1import streamlit as st2import torch3import clip4from PIL import Image5import os6import pandas as pd7from datetime import datetime8import torch.nn.functional as F9from typing import List10 11# Device setup12device = "cuda" if torch.cuda.is_available() else "cpu"13 14# Load CLIP model and preprocessor (ViT-B/32 = small model, CPU-friendly)15model, preprocess = clip.load("ViT-B/32", device=device)16model.eval()17 18# Display app title and information19st.set_page_config(page_title="Few-Shot Fault Detection", layout="wide")20st.title("๐ ๏ธ Few-Shot Fault Detection (Industrial Quality Control)")21 22st.markdown("""23This demo uses the **smaller `ViT-B/32` encoder** from OpenAI's CLIP model to classify test images as **Nominal** or **Defective**, based on few-shot learning using user-provided reference images.24 25โ ๏ธ **Note**: This app is running on a **free CPU tier** and is meant for demonstration purposes. For more advanced use cases, including GPU acceleration, custom training, and larger models, please refer to:26 27๐ [Megahed et al. (2025)](https://arxiv.org/abs/2501.12596): 28*Adapting OpenAI's CLIP Model for Few-Shot Image Inspection in Manufacturing Quality Control: An Expository Case Study with Multiple Application Examples*29 30๐ [GitHub & Colab links available in the paper](https://arxiv.org/abs/2501.12596)31""")32 33# --- Few-shot classification logic ---34def few_shot_fault_classification(35 test_images: List[Image.Image],36 test_image_filenames: List[str],37 nominal_images: List[Image.Image],38 nominal_descriptions: List[str],39 defective_images: List[Image.Image],40 defective_descriptions: List[str],41 num_few_shot_nominal_imgs: int,42 file_path: str = '.',43 file_name: str = 'image_classification_results.csv',44 print_one_liner: bool = False45):46 if not isinstance(test_images, list): test_images = [test_images]47 if not isinstance(test_image_filenames, list): test_image_filenames = [test_image_filenames]48 if not isinstance(nominal_images, list): nominal_images = [nominal_images]49 if not isinstance(nominal_descriptions, list): nominal_descriptions = [nominal_descriptions]50 if not isinstance(defective_images, list): defective_images = [defective_images]51 if not isinstance(defective_descriptions, list): defective_descriptions = [defective_descriptions]52 53 csv_file = os.path.join(file_path, file_name)54 results = []55 56 with torch.no_grad():57 nominal_features = torch.stack([model.encode_image(img.unsqueeze(0)).squeeze(0).to(device) for img in nominal_images])58 nominal_features /= nominal_features.norm(dim=-1, keepdim=True)59 60 defective_features = torch.stack([model.encode_image(img.unsqueeze(0)).squeeze(0).to(device) for img in defective_images])61 defective_features /= defective_features.norm(dim=-1, keepdim=True)62 63 csv_data = []64 65 for idx, test_img in enumerate(test_images):66 test_features = model.encode_image(test_img.unsqueeze(0)).squeeze(0).to(device)67 test_features /= test_features.norm(dim=-1, keepdim=True)68 69 max_nom_sim, max_def_sim = -float('inf'), -float('inf')70 max_nom_idx, max_def_idx = -1, -171 72 for i in range(nominal_features.shape[0]):73 sim = (test_features @ nominal_features[i].T).item()74 if sim > max_nom_sim:75 max_nom_sim, max_nom_idx = sim, i76 77 for j in range(defective_features.shape[0]):78 sim = (test_features @ defective_features[j].T).item()79 if sim > max_def_sim:80 max_def_sim, max_def_idx = sim, j81 82 similarities = torch.tensor([max_nom_sim, max_def_sim])83 probabilities = F.softmax(similarities, dim=0).tolist()84 prob_nom, prob_def = probabilities85 86 classification = "Defective" if prob_def > prob_nom else "Nominal"87 88 csv_data.append({89 "datetime_of_operation": datetime.now().isoformat(),90 "num_few_shot_nominal_imgs": num_few_shot_nominal_imgs,91 "image_path": test_image_filenames[idx],92 "image_name": test_image_filenames[idx].split('/')[-1],93 "classification_result": classification,94 "non_defect_prob": round(prob_nom, 3),95 "defect_prob": round(prob_def, 3),96 "nominal_description": nominal_descriptions[max_nom_idx],97 "defective_description": defective_descriptions[max_def_idx] if defective_images else "N/A"98 })99 100 if print_one_liner:101 print(f"{test_image_filenames[idx]} classified as {classification} "102 f"(Nominal Prob: {prob_nom:.3f}, Defective Prob: {prob_def:.3f})")103 104 file_exists = os.path.isfile(csv_file)105 with open(csv_file, mode='a' if file_exists else 'w', newline='') as file:106 import csv107 fieldnames = [108 "datetime_of_operation", "num_few_shot_nominal_imgs", "image_path", "image_name",109 "classification_result", "non_defect_prob", "defect_prob",110 "nominal_description", "defective_description"111 ]112 writer = csv.DictWriter(file, fieldnames=fieldnames)113 if not file_exists:114 writer.writeheader()115 for row in csv_data:116 writer.writerow(row)117 118 return ""119 120# --- App state ---121if 'nominal_images' not in st.session_state:122 st.session_state.nominal_images = []123if 'defective_images' not in st.session_state:124 st.session_state.defective_images = []125if 'test_images' not in st.session_state:126 st.session_state.test_images = []127if 'results' not in st.session_state:128 st.session_state.results = []129 130# --- Tabs ---131tab1, tab2, tab3 = st.tabs(["๐ฅ Upload Reference Images", "๐ Test Classification", "๐ Results"])132 133# Tab 1: Upload Reference Images134with tab1:135 st.header("Upload Reference Images")136 nominal_files = st.file_uploader("Upload Nominal Images", accept_multiple_files=True, type=['png', 'jpg', 'jpeg'])137 defective_files = st.file_uploader("Upload Defective Images", accept_multiple_files=True, type=['png', 'jpg', 'jpeg'])138 139 if nominal_files:140 st.session_state.nominal_images = [preprocess(Image.open(file).convert("RGB")).to(device) for file in nominal_files]141 st.session_state.nominal_descriptions = [file.name for file in nominal_files]142 st.success(f"Uploaded {len(nominal_files)} nominal images.")143 144 if defective_files:145 st.session_state.defective_images = [preprocess(Image.open(file).convert("RGB")).to(device) for file in defective_files]146 st.session_state.defective_descriptions = [file.name for file in defective_files]147 st.success(f"Uploaded {len(defective_files)} defective images.")148 149# Tab 2: Test Classification150with tab2:151 st.header("Upload Test Image(s)")152 test_files = st.file_uploader("Upload Test Images", accept_multiple_files=True, type=['png', 'jpg', 'jpeg'])153 154 if st.button("๐ Run Classification") and test_files:155 test_images = [preprocess(Image.open(file).convert("RGB")).to(device) for file in test_files]156 test_filenames = [file.name for file in test_files]157 158 few_shot_fault_classification(159 test_images=test_images,160 test_image_filenames=test_filenames,161 nominal_images=st.session_state.nominal_images,162 nominal_descriptions=st.session_state.nominal_descriptions,163 defective_images=st.session_state.defective_images,164 defective_descriptions=st.session_state.defective_descriptions,165 num_few_shot_nominal_imgs=len(st.session_state.nominal_images),166 file_path=".",167 file_name="streamlit_results.csv",168 print_one_liner=False169 )170 171 st.success("Classification complete!")172 st.session_state.results = "streamlit_results.csv"173 174# Tab 3: View/Download Results175with tab3:176 st.header("Classification Results")177 if os.path.exists("streamlit_results.csv"):178 df = pd.read_csv("streamlit_results.csv")179 st.dataframe(df)180 st.download_button("๐ฅ Download Results", data=df.to_csv(index=False), file_name="classification_results.csv", mime="text/csv")181 else:182 st.info("No results yet. Please classify some test images.")183 