File size: 4,648 Bytes
1a1149c
132cfbe
8f78c90
bc97233
 
8f78c90
132cfbe
7c4a59c
bc97233
 
 
 
 
 
8fcc1ae
bc97233
 
 
8fcc1ae
 
 
 
bc97233
 
 
 
8fcc1ae
 
 
 
bc97233
 
 
 
7c4a59c
8fcc1ae
 
 
 
7c4a59c
7536239
7c4a59c
 
8f78c90
 
7c4a59c
 
8f78c90
7c4a59c
8f78c90
7c4a59c
8f78c90
 
 
7c4a59c
8f78c90
bc97233
 
8f78c90
 
7c4a59c
 
8f78c90
 
 
 
 
 
 
bc97233
 
 
8f78c90
 
 
 
 
 
 
 
 
 
 
 
8fcc1ae
8f78c90
 
 
 
bc97233
8f78c90
 
 
 
 
 
 
bc97233
8f78c90
 
 
 
7c4a59c
8f78c90
7c4a59c
 
bc97233
 
 
 
 
 
 
7c4a59c
 
8fcc1ae
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
106
107
108
109
110
111
112
113
114
115
116
import gradio as gr
import spaces
from roi_predictor import predict_rois
from landmark_predictor import predict_landmark, show_landmarks
from cpak_calculator import calculate_cpak

@spaces.GPU
def predict_rois_gpu(image):
    annotated_img, gallery_items, rois = predict_rois(image)
    return annotated_img, gallery_items, rois

@spaces.GPU
def predict_landmark_gpu(image, rois):
    landmarks = predict_landmark(image, rois)
    heatmap_image, heatmap_gallery, landmark_image = show_landmarks(image, rois, landmarks)
    cpak_right_dict, cpak_left_dict = calculate_cpak(image, landmarks)
    
    right_text = (
        f"MPTA: {cpak_right_dict['MPTA']:.2f} °\n"
        f"LDFA: {cpak_right_dict['LDFA']:.2f} °\n"
        f"aHKA: {cpak_right_dict['aHKA']:.2f} ° ({cpak_right_dict['aHKA_label']})\n"
        f"JLO: {cpak_right_dict['JLO']:.2f} ° ({cpak_right_dict['JLO_label']})\n\n"
        f"CPAK: {cpak_right_dict['CPAK']} \n"
    )
    
    left_text = (
        f"MPTA: {cpak_left_dict['MPTA']:.2f} °\n"
        f"LDFA: {cpak_left_dict['LDFA']:.2f} °\n"
        f"aHKA: {cpak_left_dict['aHKA']:.2f} ° ({cpak_left_dict['aHKA_label']})\n"
        f"JLO: {cpak_left_dict['JLO']:.2f} ° ({cpak_left_dict['JLO_label']})\n\n"
        f"CPAK: {cpak_left_dict['CPAK']} \n"
    )
    
    return heatmap_image, heatmap_gallery, landmark_image, right_text, left_text

CUSTOM_CSS = """
.caption-label {color: #888;}
"""

# Define modern Gradio Blocks interface
with gr.Blocks() as demo:
    gr.Markdown(
        """
        # 🦴 Knee Landmarks ROI Detector & Analyzer
        Upload a knee X-ray image to automatically detect regions of interest (ROIs), predict landmarks, and compute CPAK measurements.
        """
    )
    
    with gr.Row():
        # Left Column: Input and Controls
        with gr.Column(scale=1):
            with gr.Tabs():
                with gr.TabItem("Input"):
                    input_image = gr.Image(type="numpy", label="Input Knee X-Ray Image")
            
            
            detect_roi_btn = gr.Button("Detect ROIs", variant="primary")
            predict_landmark_btn = gr.Button("Predict Landmarks", variant="secondary")
            
        # Right Column: Visualizations & Analysis
        with gr.Column(scale=2):
            with gr.Tabs():
                # Tab 1: ROI Detection Output
                with gr.TabItem("1. ROI Detection"):
                    with gr.Row():
                        with gr.Column(scale=1):
                            roi_image = gr.Image(label="Annotated ROIs")
                        with gr.Column(scale=1):
                            roi_gallery = gr.Gallery(label="Cropped Regions", columns=3, height="auto")

                    roi_state = gr.State()

                # Tab 2: Landmarks and Heatmaps
                with gr.TabItem("2. Landmarks"):
                    with gr.Row():
                        with gr.Column(scale=1):
                            heatmap_image = gr.Image(label="Predicted Heatmap")
                        with gr.Column(scale=1):
                            landmark_image = gr.Image(label="Final Landmarks")
                    
                    gr.Markdown("## Landmark Heatmaps")
                    heatmap_gallery = gr.Gallery(label="Heatmap Gallery", columns=4, height="auto")
                
                # Tab 3: Clinical Measurements
                with gr.TabItem("3. CPAK Measurement"):
                    with gr.Row():
                        with gr.Column():
                            gr.Markdown("# Right Side")
                            cpak_right = gr.Textbox(
                                value="",
                                show_label=False,
                                lines=8,
                                interactive=False
                            )
                        with gr.Column():
                            gr.Markdown("# Left Side")
                            cpak_left = gr.Textbox(
                                value="",
                                show_label=False,
                                lines=8,
                                interactive=False
                            )
                    
    detect_roi_btn.click(
        fn=predict_rois_gpu,
        inputs=[input_image],
        outputs=[roi_image, roi_gallery, roi_state]
    )

    predict_landmark_btn.click(
        fn=predict_landmark_gpu,
        inputs=[input_image, roi_state],
        outputs=[heatmap_image, heatmap_gallery, landmark_image, cpak_right, cpak_left]
    )

demo.launch(theme=gr.Theme.from_hub("harsh8001/comic"), css=CUSTOM_CSS, ssr_mode=False)