/
RainbowRay
/
program_engineering_lab_07
Обзор
Документация
Войти
/
RainbowRay
/
program_engineering_lab_07
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
CI/CD
Аналитика
Безопасность
main
main.py
86 строк
3 KB
Frogog
fix: flake8 errors
29 май 2026, 22:44
29 май 2026, 22:44
f695ce2
Код
Авторство
О чём код?
from PIL import Image import io from skimage import io as sio import numpy as np import torch import torch.nn.functional as F from transformers import AutoModelForImageSegmentation from torchvision.transforms.functional import normalize model = AutoModelForImageSegmentation.from_pretrained("briaai/RMBG-1.4", trust_remote_code=True) def preprocess_image(im: np.ndarray, model_input_size: list) -> torch.Tensor: if len(im.shape) < 3: im = im[:, :, np.newaxis] im_tensor = torch.tensor(im, dtype=torch.float32).permute(2, 0, 1) im_tensor = F.interpolate(torch.unsqueeze(im_tensor, 0), size=model_input_size, mode='bilinear') image = torch.divide(im_tensor, 255.0) image = normalize(image, [0.5, 0.5, 0.5], [1.0, 1.0, 1.0]) return image def postprocess_image(result: torch.Tensor, im_size: list) -> np.ndarray: result = torch.squeeze(F.interpolate(result, size=im_size, mode='bilinear'), 0) ma = torch.max(result) mi = torch.min(result) result = (result - mi) / (ma - mi) im_array = (result * 255).permute(1, 2, 0).cpu().data.numpy().astype(np.uint8) im_array = np.squeeze(im_array) return im_array device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu") model.to(device) def show_image(path: str): img = Image.open(path) img.show() # show_image("cat_walking.jpg") # show_image("girafe.jpg") # show_image("dog_fight.jpg") def prepare_input(path: str): orig_im = sio.imread(path) orig_im_size = orig_im.shape[0:2] model_input_size = [1024, 1024] image = preprocess_image(orig_im, model_input_size).to(device) result = model(image) result_image = postprocess_image(result[0][0], orig_im_size) pil_mask_im = Image.fromarray(result_image) orig_image = Image.open(path) no_bg_image = orig_image.copy() no_bg_image.putalpha(pil_mask_im) no_bg_image.show() # def io_file_input(file_object): # orig_image = file_object # orig_im = np.array(orig_image) # orig_im_size = orig_im.shape[0:2] # model_input_size = [1024, 1024] # image = preprocess_image(orig_im, model_input_size).to(device) # result = model(image) # result_image = postprocess_image(result[0][0], orig_im_size) # pil_mask_im = Image.fromarray(result_image) # orig_image = Image.open(file_object) # no_bg_image = orig_image.copy() # no_bg_image.putalpha(pil_mask_im) # return no_bg_image def io_file_input(orig_im, bytes_data): orig_im_size = orig_im.shape[0:2] model_input_size = [1024, 1024] image = preprocess_image(orig_im, model_input_size).to(device) result = model(image) result_image = postprocess_image(result[0][0], orig_im_size) pil_mask_im = Image.fromarray(result_image) orig_image = Image.open(io.BytesIO(bytes_data)) no_bg_image = orig_image.copy() no_bg_image.putalpha(pil_mask_im) return no_bg_image