File size: 6,588 Bytes
625c0a2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b0cc1ff
625c0a2
 
 
 
 
8999815
 
625c0a2
8999815
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
"""
QuickDraw Classifier API
A simple API wrapper for the CLIP-based drawing classifier
"""

import requests
import base64
import json
from typing import List, Dict, Optional
from PIL import Image
import io

class QuickDrawAPI:
    """
    API client for the QuickDraw Classifier
    """
    
    def __init__(self, base_url: str):
        """
        Initialize the API client
        
        Args:
            base_url: Base URL of the deployed Hugging Face Space
                     (e.g., "https://huggingface.co/spaces/username/quickdraw-classifier")
        """
        self.base_url = base_url.rstrip('/')
    
    def classify_image_file(self, image_path: str, top_k: int = 5) -> Dict:
        """
        Classify a drawing from an image file
        
        Args:
            image_path: Path to the image file
            top_k: Number of top predictions to return
            
        Returns:
            Dictionary with classification results
        """
        with open(image_path, "rb") as f:
            image_data = base64.b64encode(f.read()).decode()
        
        return self.classify_image_base64(image_data, top_k)
    
    def classify_pil_image(self, image: Image.Image, top_k: int = 5) -> Dict:
        """
        Classify a PIL Image
        
        Args:
            image: PIL Image object
            top_k: Number of top predictions to return
            
        Returns:
            Dictionary with classification results
        """
        # Convert PIL image to base64
        buffer = io.BytesIO()
        image.save(buffer, format='PNG')
        image_data = base64.b64encode(buffer.getvalue()).decode()
        
        return self.classify_image_base64(image_data, top_k)
    
    def classify_image_base64(self, image_data: str, top_k: int = 5) -> Dict:
        """
        Classify a base64 encoded image
        
        Args:
            image_data: Base64 encoded image string
            top_k: Number of top predictions to return
            
        Returns:
            Dictionary with classification results
        """
        try:
            response = requests.post(
                f"{self.base_url}/api/predict",
                json={
                    "data": [image_data, top_k],
                    "fn_index": 0
                },
                timeout=30
            )
            
            if response.status_code == 200:
                result = response.json()
                # Parse Gradio response format
                if "data" in result and len(result["data"]) > 0:
                    return {
                        "success": True,
                        "predictions": self._parse_gradio_output(result["data"][0])
                    }
            
            return {
                "success": False,
                "error": f"API request failed with status {response.status_code}"
            }
            
        except Exception as e:
            return {
                "success": False,
                "error": str(e)
            }
    
    def _parse_gradio_output(self, output: str) -> List[Dict]:
        """
        Parse the Gradio markdown output to extract predictions
        
        Args:
            output: Markdown formatted output from Gradio
            
        Returns:
            List of prediction dictionaries
        """
        predictions = []
        
        # Simple parsing of the markdown output
        lines = output.split('\n')
        for line in lines:
            if line.strip() and any(char.isdigit() for char in line):
                # Look for lines like "1. **Cat** - 85.2%"
                parts = line.split('-')
                if len(parts) >= 2:
                    # Extract category name
                    left_part = parts[0].strip()
                    category_start = left_part.find('**') + 2
                    category_end = left_part.rfind('**')
                    
                    if category_start > 1 and category_end > category_start:
                        category = left_part[category_start:category_end].strip().lower()
                        
                        # Extract confidence
                        right_part = parts[1].strip()
                        confidence_str = right_part.replace('%', '').strip()
                        try:
                            confidence = float(confidence_str) / 100.0
                            predictions.append({
                                "category": category,
                                "confidence": confidence
                            })
                        except ValueError:
                            continue
        
        return predictions

# Example usage functions
def classify_drawing_simple(image_path: str, space_url: str) -> List[str]:
    """
    Simple function to get top categories for a drawing
    
    Args:
        image_path: Path to the drawing image
        space_url: URL of the deployed Hugging Face Space
        
    Returns:
        List of top category names
    """
    api = QuickDrawAPI(space_url)
    result = api.classify_image_file(image_path)
    
    if result["success"]:
        return [pred["category"] for pred in result["predictions"]]
    else:
        print(f"Error: {result['error']}")
        return []

def batch_classify_drawings(image_paths: List[str], space_url: str) -> Dict[str, List[str]]:
    """
    Classify multiple drawings at once
    
    Args:
        image_paths: List of paths to drawing images
        space_url: URL of the deployed Hugging Face Space
        
    Returns:
        Dictionary mapping image paths to predicted categories
    """
    api = QuickDrawAPI(space_url)
    results = {}
    
    for image_path in image_paths:
        categories = classify_drawing_simple(image_path, space_url)
        results[image_path] = categories
    
    return results

# Example usage
if __name__ == "__main__":
    # Example usage of the API
    SPACE_URL = "https://huggingface.co/spaces/souvikg544/quickdraw-classifier"  # Replace with your space URL
    
    # Initialize API client
    api = QuickDrawAPI(SPACE_URL)
    
    # Example: Classify an image file
    result = api.classify_image_file("temp_drawing.png")
    print(json.dumps(result, indent=2))
    
    # print("QuickDraw API client ready!")
    # print(f"Connect to your space at: {SPACE_URL}")
    # print("\nExample usage:")
    # print("api = QuickDrawAPI('https://your-space-url')")
    # print("result = api.classify_image_file('drawing.png')")
    # print("print(result['predictions'])")