Spaces:
Sleeping
Sleeping
File size: 4,599 Bytes
f19f69c | 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 | 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 |