VVV-Chain-Model / inference.py
LouayMagdy's picture
Upload folder using huggingface_hub
f19f69c verified
Raw
History Blame Contribute Delete
4.6 kB
import torch
import tempfile
import numpy as np
import nibabel as nib
from fastapi import FastAPI, UploadFile, File
def preprocess_input(flair_file: UploadFile, t1_file: UploadFile, t1ce_file: UploadFile, t2_file: UploadFile):
with tempfile.NamedTemporaryFile(suffix=".nii.gz") as temp_flair, \
tempfile.NamedTemporaryFile(suffix=".nii.gz") as temp_t1, \
tempfile.NamedTemporaryFile(suffix=".nii.gz") as temp_t1ce, \
tempfile.NamedTemporaryFile(suffix=".nii.gz") as temp_t2:
# Save uploaded files to temporary NIfTI files
flair_content = flair_file.file.read()
t1_content = t1_file.file.read()
t1ce_content = t1ce_file.file.read()
t2_content = t2_file.file.read()
if flair_file.filename.endswith('.gz'):
with gzip.open(temp_flair.name, 'wb') as f:
f.write(flair_content)
else:
with open(temp_flair.name, 'wb') as f:
f.write(flair_content)
if t1_file.filename.endswith('.gz'):
with gzip.open(temp_t1.name, 'wb') as f:
f.write(t1_content)
else:
with open(temp_t1.name, 'wb') as f:
f.write(t1_content)
if t1ce_file.filename.endswith('.gz'):
with gzip.open(temp_t1ce.name, 'wb') as f:
f.write(t1ce_content)
else:
with open(temp_t1ce.name, 'wb') as f:
f.write(t1ce_content)
if t2_file.filename.endswith('.gz'):
with gzip.open(temp_t2.name, 'wb') as f:
f.write(t2_content)
else:
with open(temp_t2.name, 'wb') as f:
f.write(t2_content)
# Load and preprocess the NIfTI files
flair = nib.load(temp_flair.name).get_fdata()
t1 = nib.load(temp_t1.name).get_fdata()
t1ce = nib.load(temp_t1ce.name).get_fdata()
t2 = nib.load(temp_t2.name).get_fdata()
# Process the arrays
flair_cropped = flair[56:184, 56:184, 13:141]
t1_cropped = t1[56:184, 56:184, 13:141]
t1ce_cropped = t1ce[56:184, 56:184, 13:141]
t2_cropped = t2[56:184, 56:184, 13:141]
# Convert to PyTorch tensors
flair_tensor = torch.tensor(flair_cropped, dtype=torch.float32).unsqueeze(0)
t1_tensor = torch.tensor(t1_cropped, dtype=torch.float32).unsqueeze(0)
t1ce_tensor = torch.tensor(t1ce_cropped, dtype=torch.float32).unsqueeze(0)
t2_tensor = torch.tensor(t2_cropped, dtype=torch.float32).unsqueeze(0)
return flair_tensor, t1_tensor, t1ce_tensor, t2_tensor
def segment_wt(wt_model, flair_tensor, t1_tensor, t1ce_tensor, t2_tensor):
wt_input = torch.cat((flair_tensor, t1_tensor, t1ce_tensor, t2_tensor), dim=0).unsqueeze(0)
with torch.no_grad():
wt_output = wt_model(wt_input)
return wt_output
def segment_tc(tc_model, flair_tensor, t1_tensor, t1ce_tensor, t2_tensor, wt_output):
tc_input = torch.cat((flair_tensor, t1_tensor, t1ce_tensor, t2_tensor, wt_output.squeeze(0),
wt_output.squeeze(0), wt_output.squeeze(0), wt_output.squeeze(0)), dim=0).unsqueeze(0)
with torch.no_grad():
tc_output = tc_model(tc_input)
return tc_output
def segment_et(et_model, flair_tensor, t1_tensor, t1ce_tensor, t2_tensor, wt_output, tc_output):
et_input = torch.cat((flair_tensor, t1_tensor, t1ce_tensor, t2_tensor, wt_output.squeeze(0),
wt_output.squeeze(0), tc_output.squeeze(0), tc_output.squeeze(0)), dim=0).unsqueeze(0)
with torch.no_grad():
et_output = et_model(et_input)
return et_output
def segment_all(wt_model, tc_model, et_model, flair_tensor, t1_tensor, t1ce_tensor, t2_tensor):
wt_output = segment_wt(wt_model, flair_tensor, t1_tensor, t1ce_tensor, t2_tensor)
tc_output = segment_tc(tc_model, flair_tensor, t1_tensor, t1ce_tensor, t2_tensor, wt_output)
et_output = segment_et(et_model, flair_tensor, t1_tensor, t1ce_tensor, t2_tensor, wt_output, tc_output)
wt_label = torch.sigmoid(wt_output.squeeze(0).squeeze(0)) > 0.5
tc_label = torch.sigmoid(tc_output.squeeze(0).squeeze(0)) > 0.5
et_label = torch.sigmoid(et_output.squeeze(0).squeeze(0)) > 0.5
output = np.zeros((128, 128, 128))
output[(tc_label == 1) & (wt_label == 0)] = 2
output[(et_label == 1) & (tc_label == 0)] = 3
output[(et_label == 0) & (tc_label == 0)] = 1
output_in_original_size = np.zeros((240, 240, 155))
output_in_original_size[56:184, 56:184, 13:141] = output
return output_in_original_size