CoolFace
Apppublic

fmegahed/clip

sourceHugging Facemitupdated 1y agoView on Hugging Face
1likes
app.py183 linesDownload Raw Back to root
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