Team Ai
Apppublic

ankitkr1499/Image_To_Sketch_With_Content_Based_Feature_Extraction

sourceHugging Faceupdated 1y agoView on Hugging Face
4likes
utils.py164 linesDownload Raw Back to root
1import torch
2from PIL import Image
3import cv2
4import numpy as np
5import torchvision.transforms as transforms
6
7
8from transformers import BlipProcessor, BlipForConditionalGeneration
9from PIL import Image
10import torch
11
12# Setup device
13device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
14
15# Load BLIP model and processor
16blip_processor = BlipProcessor.from_pretrained("Salesforce/blip-image-captioning-base")
17blip_model = BlipForConditionalGeneration.from_pretrained("Salesforce/blip-image-captioning-base").to(device)
18
19def get_image_caption(image: Image.Image) -> str:
20    image = image.convert("RGB")
21    inputs = blip_processor(images=image, return_tensors="pt").to(device)
22    output = blip_model.generate(**inputs)
23    caption = blip_processor.decode(output[0], skip_special_tokens=True)
24    return caption
25
26
27def generate_image_from_sketch(sketch_img: Image.Image, generator_model: torch.nn.Module, device='cpu') -> Image.Image:
28    # Resize and convert to RGB to ensure 3 channels
29    sketch_img = sketch_img.resize((256, 256)).convert("RGB")
30
31    # Convert to tensor and normalize to [-1, 1]
32    transform = transforms.Compose([
33        transforms.ToTensor(),  # [0, 1]
34        transforms.Normalize([0.5]*3, [0.5]*3)  # → [-1, 1]
35    ])
36    input_tensor = transform(sketch_img).unsqueeze(0).to(device)
37
38    # Generate output
39    generator_model.to(device)
40    generator_model.eval()
41    with torch.no_grad():
42        output_tensor = generator_model(input_tensor)
43
44    # Denormalize output back to [0, 1]
45    output_tensor = output_tensor.squeeze(0).cpu() * 0.5 + 0.5
46    output_img = transforms.ToPILImage()(output_tensor.clamp(0, 1))
47
48    return output_img
49
50def dodge_sketch(image: Image.Image) -> Image.Image:
51    image_np = np.array(image.convert("RGB"))
52    gray = cv2.cvtColor(image_np, cv2.COLOR_RGB2GRAY)
53    inverted = 255 - gray
54    blur = cv2.GaussianBlur(inverted, (21, 21), sigmaX=0, sigmaY=0)
55    dodge = cv2.divide(gray, 255 - blur, scale=256)
56    edges = cv2.Laplacian(gray, cv2.CV_8U, ksize=5)
57    edges = cv2.GaussianBlur(edges, (3, 3), 0)
58    sketch = cv2.subtract(dodge, edges // 3)
59    sketch = np.clip(sketch, 0, 255).astype(np.uint8)
60    sketch = cv2.fastNlMeansDenoising(sketch, h=15, templateWindowSize=7, searchWindowSize=21)
61    return Image.fromarray(sketch).convert("RGB")
62
63def sobel_sketch(image: Image.Image) -> Image.Image:
64    # Convert PIL Image to NumPy array (RGB) and then to grayscale
65    image_np = np.array(image.convert("RGB"))
66    gray = cv2.cvtColor(image_np, cv2.COLOR_RGB2GRAY)
67    sobel_x = cv2.Sobel(gray, cv2.CV_64F, 1, 0, ksize=3)
68    sobel_y = cv2.Sobel(gray, cv2.CV_64F, 0, 1, ksize=3)
69    
70    abs_sobel_x = cv2.convertScaleAbs(sobel_x)
71    abs_sobel_y = cv2.convertScaleAbs(sobel_y)
72    
73    sobel_combined = cv2.addWeighted(abs_sobel_x, 0.5, abs_sobel_y, 0.5, 0)
74
75    inverted_edges = 255 - sobel_combined
76    
77    sketch = cv2.fastNlMeansDenoising(inverted_edges, h=15, templateWindowSize=7, searchWindowSize=21)
78    
79    return Image.fromarray(sketch).convert("RGB")
80
81def lattice_sketch(image: Image.Image) -> Image.Image:
82    image_np = np.array(image.convert("RGB"))
83    gray = cv2.cvtColor(image_np, cv2.COLOR_RGB2GRAY)
84    sharpen_kernel = np.array([[0, -1, 0],
85                               [-1, 5, -1],
86                               [0, -1, 0]])
87    sharpened = cv2.filter2D(gray, -1, sharpen_kernel)
88    inverted = 255 - sharpened
89    blur = cv2.GaussianBlur(inverted, (21, 21), 0)
90    dodge = cv2.divide(sharpened, 255 - blur, scale=256)
91    sketch = cv2.equalizeHist(dodge)
92    sketch = cv2.fastNlMeansDenoising(sketch, h=75, templateWindowSize=7, searchWindowSize=21)
93    return Image.fromarray(sketch).convert("RGB")
94
95def canny_sketch(image: Image.Image) -> Image.Image:
96    def gaussian_blur(image, kernel_size=5, sigma=1.4):
97        return cv2.GaussianBlur(image, (kernel_size, kernel_size), sigma)
98
99    def sobel_filters(img):
100        Ix = cv2.Sobel(img, cv2.CV_64F, 1, 0, ksize=3)
101        Iy = cv2.Sobel(img, cv2.CV_64F, 0, 1, ksize=3)
102        magnitude = np.hypot(Ix, Iy)
103        angle = np.arctan2(Iy, Ix)
104        magnitude = magnitude / magnitude.max() * 255
105        return magnitude.astype(np.uint8), angle
106
107    def non_max_suppression(mag, angle):
108        M, N = mag.shape
109        output = np.zeros((M, N), dtype=np.uint8)
110        angle = angle * 180 / np.pi
111        angle[angle < 0] += 180
112        for i in range(1, M - 1):
113            for j in range(1, N - 1):
114                q, r = 255, 255
115                a = angle[i, j]
116                if (0 <= a < 22.5) or (157.5 <= a <= 180):
117                    q = mag[i, j + 1]
118                    r = mag[i, j - 1]
119                elif (22.5 <= a < 67.5):
120                    q = mag[i + 1, j - 1]
121                    r = mag[i - 1, j + 1]
122                elif (67.5 <= a < 112.5):
123                    q = mag[i + 1, j]
124                    r = mag[i - 1, j]
125                elif (112.5 <= a < 157.5):
126                    q = mag[i - 1, j - 1]
127                    r = mag[i + 1, j + 1]
128                output[i, j] = mag[i, j] if mag[i, j] >= q and mag[i, j] >= r else 0
129        return output
130
131    def threshold(image, low_ratio=0.05, high_ratio=0.15):
132        high_threshold = image.max() * high_ratio
133        low_threshold = high_threshold * low_ratio
134        M, N = image.shape
135        res = np.zeros((M, N), dtype=np.uint8)
136        strong, weak = np.uint8(255), np.uint8(75)
137        strong_i, strong_j = np.where(image >= high_threshold)
138        weak_i, weak_j = np.where((image <= high_threshold) & (image >= low_threshold))
139        res[strong_i, strong_j] = strong
140        res[weak_i, weak_j] = weak
141        return res, weak, strong
142
143    def hysteresis(img, weak=75, strong=255):
144        M, N = img.shape
145        for i in range(1, M - 1):
146            for j in range(1, N - 1):
147                if img[i, j] == weak:
148                    if ((img[i+1, j-1] == strong) or (img[i+1, j] == strong) or (img[i+1, j+1] == strong)
149                        or (img[i, j-1] == strong) or (img[i, j+1] == strong)
150                        or (img[i-1, j-1] == strong) or (img[i-1, j] == strong) or (img[i-1, j+1] == strong)):
151                        img[i, j] = strong
152                    else:
153                        img[i, j] = 0
154        return img
155
156    image_np = np.array(image.convert("RGB"))
157    gray = cv2.cvtColor(image_np, cv2.COLOR_RGB2GRAY)
158    blurred = gaussian_blur(gray)
159    mag, angle = sobel_filters(blurred)
160    nms = non_max_suppression(mag, angle)
161    thresh, weak, strong = threshold(nms)
162    result = hysteresis(thresh, weak, strong)
163    inverted = 255 - result
164    return Image.fromarray(inverted).convert("RGB")