import gc import os from pathlib import Path from typing import Dict, List, Union from concurrent.futures import ThreadPoolExecutor, as_completed import torch import torch.nn as nn from PIL import Image from transformers import CLIPVisionModel, CLIPProcessor from peft import LoraConfig, get_peft_model import warnings warnings.filterwarnings(action="ignore") class CONSTANTS: BASE_MODEL_NAME = "openai/clip-vit-base-patch16" TUNED_MODEL_NAME = "industrial-document-classifier-clip-lora.pt" EMBEDDING_DIM = 768 NUM_PARENT_CLASSES = 6 NUM_CHILD_CLASSES = 13 DEVICE = "cuda" if torch.cuda.is_available() else "cpu" LORA_R = 32 LORA_ALPHA = 64 LORA_DROPOUT = 0.1 LORA_TARGET_MODULES = ["q_proj", "k_proj", "v_proj", "out_proj", "fc1", "fc2"] BATCH_SIZE = 16 MAX_WORKERS = os.cpu_count() VALID_IMAGE_EXTENSIONS = {'.jpg', '.jpeg', '.png', '.bmp', '.gif', '.tiff', '.tif', '.webp', '.ico', '.heic', '.heif'} PARENT_CLASS_NAMES = {0: "product_information", 1: "engineering_drawings", 2: "instructional_guides", 3: "compliance_certificates", 4: "energy_ratings", 5: "warranty_documents"} class HierarchicalDocumentClassifier(nn.Module): def __init__(self, model_name: str, num_parent_classes: int, num_child_classes: int, embedding_dim: int): super().__init__() self.vision_model = CLIPVisionModel.from_pretrained(model_name, use_safetensors=False) self.parent_classifier = nn.Linear(embedding_dim, num_parent_classes) self.child_classifier = nn.Linear(embedding_dim, num_child_classes) def forward(self, pixel_values): outputs = self.vision_model(pixel_values=pixel_values) embeddings = outputs.pooler_output parent_logits = self.parent_classifier(embeddings) child_logits = self.child_classifier(embeddings) return parent_logits, child_logits, embeddings class DocumentClassifier: def __init__(self): self.device = torch.device(CONSTANTS.DEVICE) self.processor = CLIPProcessor.from_pretrained(CONSTANTS.BASE_MODEL_NAME, use_fast=True) self.model = self._load_model() self.model.eval() self._clear_cache() def _load_model(self) -> nn.Module: model = HierarchicalDocumentClassifier( model_name=CONSTANTS.BASE_MODEL_NAME, num_parent_classes=CONSTANTS.NUM_PARENT_CLASSES, num_child_classes=CONSTANTS.NUM_CHILD_CLASSES, embedding_dim=CONSTANTS.EMBEDDING_DIM) lora_config = LoraConfig( r=CONSTANTS.LORA_R, lora_alpha=CONSTANTS.LORA_ALPHA, target_modules=CONSTANTS.LORA_TARGET_MODULES, lora_dropout=CONSTANTS.LORA_DROPOUT, bias="none") model.vision_model = get_peft_model(model.vision_model, lora_config) checkpoint = torch.load(CONSTANTS.TUNED_MODEL_NAME, map_location=self.device) model.load_state_dict(checkpoint['model_state_dict']) del checkpoint self._clear_cache() return model.to(self.device) def _clear_cache(self): if torch.cuda.is_available(): torch.cuda.empty_cache() gc.collect() def _is_valid_image(self, image_path: str) -> bool: path = Path(image_path) if not path.exists(): return False if not path.is_file(): return False if path.suffix.lower() not in CONSTANTS.VALID_IMAGE_EXTENSIONS: return False return True def _load_image(self, image_path: str) -> Union[Image.Image, None]: try: img = Image.open(image_path).convert('RGB') return img except Exception: return None def _process_batch(self, images: List[Image.Image]): try: inputs = self.processor(images=images, return_tensors="pt") pixel_values = inputs['pixel_values'].to(self.device) with torch.no_grad(): parent_logits, _, _ = self.model(pixel_values) probabilities = torch.softmax(parent_logits, dim=1) return probabilities.cpu() except Exception: return None def _process_single_image(self, image_path: str) -> Dict: result = { "image_path": image_path, "predictions": {}, "error_response": "" } if not self._is_valid_image(image_path): result["error_response"] = "Invalid image path or unsupported format" return result img = self._load_image(image_path) if img is None: result["error_response"] = "Failed to load image" return result probabilities = self._process_batch([img]) if probabilities is None: result["error_response"] = "Model inference failed" return result probs = probabilities[0] predictions_dict = {} for class_id, class_name in CONSTANTS.PARENT_CLASS_NAMES.items(): predictions_dict[class_name] = float(probs[class_id]) result["predictions"] = predictions_dict return result def _process_batch_images(self, image_paths: List[str]) -> List[Dict]: batch_results = [] batch_images = [] batch_valid_paths = [] for img_path in image_paths: result = { "image_path": img_path, "predictions": {}, "error_response": "" } if not self._is_valid_image(img_path): result["error_response"] = "Invalid image path or unsupported format" batch_results.append(result) continue img = self._load_image(img_path) if img is None: result["error_response"] = "Failed to load image" batch_results.append(result) continue batch_images.append(img) batch_valid_paths.append(img_path) batch_results.append(result) if len(batch_images) > 0: probabilities = self._process_batch(batch_images) if probabilities is None: for result in batch_results: if result["error_response"] == "": result["error_response"] = "Model inference failed" else: valid_idx = 0 for result in batch_results: if result["error_response"] == "": probs = probabilities[valid_idx] predictions_dict = {} for class_id, class_name in CONSTANTS.PARENT_CLASS_NAMES.items(): predictions_dict[class_name] = float(probs[class_id]) result["predictions"] = predictions_dict valid_idx += 1 return batch_results def predict(self, image_paths: List[str]) -> List[Dict]: if len(image_paths) == 0: return [] batches = [] for i in range(0, len(image_paths), CONSTANTS.BATCH_SIZE): batches.append(image_paths[i:i + CONSTANTS.BATCH_SIZE]) results_map = {} with ThreadPoolExecutor(max_workers=CONSTANTS.MAX_WORKERS) as executor: future_to_batch = {executor.submit(self._process_batch_images, batch): batch for batch in batches} for future in as_completed(future_to_batch): batch = future_to_batch[future] batch_results = future.result() for result in batch_results: results_map[result["image_path"]] = result ordered_results = [] for img_path in image_paths: ordered_results.append(results_map[img_path]) return ordered_results def unload(self): if hasattr(self, 'model') and self.model is not None: del self.model if hasattr(self, 'processor') and self.processor is not None: del self.processor self._clear_cache()