VVV-Chain-Model / main.py
LouayMagdy's picture
Upload main.py with huggingface_hub
1b645ac verified
Raw
History Blame Contribute Delete
2.05 kB
from model_loader import load_vChain_model
from inference import preprocess_input, segment_all
import gzip
import tempfile
import numpy as np
import nibabel as nib
from fastapi import FastAPI, UploadFile, File, HTTPException, status
from fastapi.responses import FileResponse
app = FastAPI()
wt_model, tc_model, et_model = load_vChain_model('./chkpts/wt_chkpt.pth', './chkpts/tc_chkpt.pth', './chkpts/et_chkpt.pth')
is_ready = True
print("Model Started Successfully! ")
@app.get("/status/")
def status():
return {"msg": "The Model is Idle!"} if is_ready else {"msg": "The Model is Processing Some Input Now, Come On Later!"}
@app.post("/segment/")
def segment(files: list[UploadFile]):
global is_ready
try:
is_ready = False
flair, t1, t1ce, t2 = (f for f in files)
flair_tensor, t1_tensor, t1ce_tensor, t2_tensor = preprocess_input(flair, t1, t1ce, t2)
print("Input Files Processed Successfully!")
output = segment_all(wt_model, tc_model, et_model, flair_tensor, t1_tensor, t1ce_tensor, t2_tensor)
print("Segmentated Successfuly!\nSending Output File ...")
output_img = nib.Nifti1Image(output, affine=np.eye(4))
with tempfile.NamedTemporaryFile(suffix=".nii", delete=False) as tempFile:
nib.save(output_img, tempFile.name)
tempFile.seek(0)
with tempfile.NamedTemporaryFile(suffix=".nii.gz", delete=False) as tempGzipFile:
with gzip.open(tempGzipFile.name, 'wb') as gz:
gz.write(tempFile.read())
response = FileResponse(tempGzipFile.name, media_type="multipart/form-data")
tempGzipFile.close()
tempFile.close()
is_ready = True
return response
except Exception as e:
is_ready = True
print(f"Error occurred: {str(e)}")
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="An error occurred during segmentation.")