Team Ai
Apppublic

JayyyyAWW/MRI_Preprocessing

sourceHugging Faceupdated 10mo agoView on Hugging Face
0likes
app.py811 linesDownload Raw Back to root
1"""
2Gradio app for MRI Data Preprocessing Pipeline
3Supports bias field correction, head mask creation, and image visualization
4"""
5
6import os
7import tempfile
8import gradio as gr
9import numpy as np
10import SimpleITK as sitk
11import matplotlib.pyplot as plt
12from matplotlib.figure import Figure
13import io
14from PIL import Image
15import traceback
16from pathlib import Path
17
18# ==================== Processing Functions ====================
19
20def preprocess_mri(
21    uploaded_file,
22    bias_correction: bool = True,
23    mask_creation: bool = True,
24    intensity_normalization: bool = True,
25    slice_to_display: int = None
26):
27    """
28    Main preprocessing pipeline for MRI data
29    """
30    try:
31        # Read the uploaded file
32        if uploaded_file is None:
33            return None, None, None, "Error: Please upload an MRI image file"
34        
35        # Get file path
36        file_path = uploaded_file.name
37        
38        # Read with SimpleITK for detailed processing
39        raw_img_sitk = sitk.ReadImage(file_path, sitk.sitkFloat32)
40        raw_img_sitk = sitk.DICOMOrient(raw_img_sitk, 'RPS')
41        raw_img_arr = sitk.GetArrayFromImage(raw_img_sitk)
42        
43        # Set default slice if not provided
44        if slice_to_display is None:
45            slice_to_display = raw_img_arr.shape[0] // 2
46        
47        # Ensure slice is within bounds
48        slice_to_display = min(slice_to_display, raw_img_arr.shape[0] - 1)
49        
50        processing_log = ["Starting MRI preprocessing pipeline..."]
51        processed_img = raw_img_sitk
52        
53        # Step 1: Intensity Normalization
54        if intensity_normalization:
55            processed_img = sitk.RescaleIntensity(processed_img, 0, 255)
56            processing_log.append("✓ Intensity normalization applied")
57        
58        # Step 2: Head mask creation
59        if mask_creation:
60            transformed = sitk.RescaleIntensity(processed_img, 0, 255)
61            head_mask = sitk.LiThreshold(transformed, 0, 1)
62            processing_log.append("✓ Head mask created using Li thresholding")
63        
64        # Step 3: Bias field correction
65        if bias_correction:
66            try:
67                shrink_factor_bias = 4
68                inputImage = sitk.Shrink(processed_img, [shrink_factor_bias] * processed_img.GetDimension())
69                maskImage = sitk.Shrink(head_mask, [shrink_factor_bias] * processed_img.GetDimension())
70                
71                bias_corrector = sitk.N4BiasFieldCorrectionImageFilter()
72                corrected = bias_corrector.Execute(inputImage, maskImage)
73                
74                log_bias_field = bias_corrector.GetLogBiasFieldAsImage(processed_img)
75                corrected_image = processed_img / sitk.Exp(log_bias_field)
76                processed_img = corrected_image
77                processing_log.append("✓ Bias field correction applied (N4)")
78            except Exception as e:
79                processing_log.append(f"⚠ Bias correction skipped: {str(e)}")
80        
81        # Prepare output visualization
82        processed_arr = sitk.GetArrayFromImage(processed_img)
83        
84        # Create comparison visualization
85        output_img = create_slice_visualization(raw_img_arr, processed_arr, slice_to_display)
86        
87        # Log message
88        log_text = "\n".join(processing_log)
89        
90        return output_img, log_text, processed_arr, "Processing completed successfully!"
91    
92    except Exception as e:
93        error_msg = f"Error: {str(e)}\n\n{traceback.format_exc()}"
94        return None, None, None, error_msg
95
96
97def create_slice_visualization(before, after, slice_idx):
98    """Create a side-by-side comparison of before and after slices"""
99    fig = plt.figure(figsize=(12, 5))
100    
101    # Ensure slice index is valid
102    slice_idx = min(slice_idx, before.shape[0] - 1)
103    
104    ax1 = fig.add_subplot(121)
105    ax1.imshow(before[slice_idx, :, :], cmap='gray')
106    ax1.set_title(f'Original (Slice {slice_idx})', fontsize=12, fontweight='bold')
107    ax1.axis('off')
108    
109    ax2 = fig.add_subplot(122)
110    ax2.imshow(after[slice_idx, :, :], cmap='gray')
111    ax2.set_title(f'Processed (Slice {slice_idx})', fontsize=12, fontweight='bold')
112    ax2.axis('off')
113    
114    plt.tight_layout()
115    
116    # Convert to image
117    buf = io.BytesIO()
118    plt.savefig(buf, format='png', dpi=100, bbox_inches='tight')
119    buf.seek(0)
120    img = Image.open(buf)
121    plt.close()
122    
123    return img
124
125
126def save_processed_image(processed_array, original_path):
127    """Save the processed image as NIfTI file"""
128    try:
129        if processed_array is None:
130            return None, "No processed image to save"
131        
132        # Read original to get metadata
133        original_img = sitk.ReadImage(original_path, sitk.sitkFloat32)
134        
135        # Create new image from processed array
136        processed_img = sitk.GetImageFromArray(processed_array)
137        
138        # Copy metadata
139        processed_img.SetSpacing(original_img.GetSpacing())
140        processed_img.SetOrigin(original_img.GetOrigin())
141        processed_img.SetDirection(original_img.GetDirection())
142        
143        # Save to temporary file
144        output_path = tempfile.NamedTemporaryFile(suffix='.nii.gz', delete=False).name
145        sitk.WriteImage(processed_img, output_path)
146        
147        return output_path, f"✓ Image saved to {output_path}"
148    
149    except Exception as e:
150        return None, f"Error saving image: {str(e)}"
151
152
153# Custom CSS for polished UI
154custom_css = """
155/* Root color variables */
156:root {
157    --bg-primary: #1a1d29;
158    --bg-secondary: #252837;
159    --bg-tertiary: #1e2130;
160    --text-primary: #ffffff;
161    --text-secondary: #e8e9f3;
162    --text-tertiary: #b4b7c9;
163    --text-muted: #9ca3bc;
164    --text-dim: #6b7280;
165    --accent: #6366f1;
166    --accent-hover: #8b5cf6;
167    --success: #10b981;
168    --border-light: #3a3f52;
169}
170
171/* Import fonts */
172@import url('https://fonts.googleapis.com/css2?family=Space+Grotesk:wght@400;500;600;700&family=Inter:wght@400;500;600;700&family=JetBrains+Mono:wght@400;500&display=swap');
173
174/* Global styles */
175* {
176    font-family: 'Inter', sans-serif;
177}
178
179body, .gradio-container {
180    background-color: var(--bg-primary) !important;
181}
182
183.gradio-container {
184    max-width: 100% !important;
185    padding: 40px !important;
186}
187
188/* Typography */
189h1 {
190    font-family: 'Space Grotesk', sans-serif;
191    font-size: 32px;
192    font-weight: 700;
193    letter-spacing: -0.5px;
194    color: var(--text-primary) !important;
195    margin-bottom: 8px;
196}
197
198h2 {
199    font-family: 'Space Grotesk', sans-serif;
200    font-size: 20px;
201    font-weight: 600;
202    color: var(--text-secondary) !important;
203}
204
205label, .label-wrap label {
206    font-size: 16px !important;
207    font-weight: 500 !important;
208    color: var(--text-tertiary) !important;
209    letter-spacing: 0;
210}
211
212.info-text {
213    font-size: 14px;
214    font-weight: 400;
215    line-height: 1.5;
216    color: var(--text-muted) !important;
217}
218
219/* Cards */
220.block {
221    border: 1px solid var(--border-light) !important;
222    border-radius: 12px !important;
223    background: var(--bg-secondary) !important;
224    box-shadow: 0 4px 6px -1px rgba(0, 0, 0, 0.1), 0 2px 4px -1px rgba(0, 0, 0, 0.06) !important;
225    padding: 32px !important;
226    transition: all 0.3s ease;
227}
228
229.block:hover {
230    border-color: rgba(99, 102, 241, 0.3) !important;
231    box-shadow: 0 10px 15px -3px rgba(99, 102, 241, 0.1), 0 4px 6px -2px rgba(0, 0, 0, 0.05) !important;
232}
233
234/* Input/Output sections */
235.input-container, .output-container {
236    background: var(--bg-secondary) !important;
237    border: 1px solid var(--border-light) !important;
238    border-radius: 12px !important;
239    padding: 32px !important;
240    margin-bottom: 40px !important;
241}
242
243/* File upload button */
244.file-upload input {
245    display: none;
246}
247
248.file-upload button, input[type="file"] + label, button.file-upload {
249    background: linear-gradient(135deg, var(--accent) 0%, var(--accent-hover) 100%) !important;
250    color: var(--text-primary) !important;
251    border: none !important;
252    border-radius: 14px !important;
253    padding: 12px 24px !important;
254    font-size: 14px !important;
255    font-weight: 500 !important;
256    cursor: pointer;
257    transition: all 0.2s ease;
258}
259
260.file-upload button:hover, input[type="file"] + label:hover {
261    transform: scale(1.02);
262    box-shadow: 0 10px 20px rgba(99, 102, 241, 0.3) !important;
263}
264
265.file-upload button:active, input[type="file"] + label:active {
266    transform: scale(0.98);
267}
268
269/* Checkboxes */
270input[type="checkbox"] {
271    width: 20px !important;
272    height: 20px !important;
273    cursor: pointer;
274    accent-color: var(--accent) !important;
275    transition: all 0.2s ease;
276}
277
278input[type="checkbox"]:hover {
279    transform: scale(1.1);
280}
281
282input[type="checkbox"]:focus-visible {
283    outline: 2px solid var(--accent) !important;
284    outline-offset: 2px;
285}
286
287.checkbox-wrap {
288    display: flex;
289    align-items: center;
290    margin: 16px 0;
291    gap: 8px;
292}
293
294.checkbox-wrap label {
295    margin: 0 !important;
296    font-size: 16px !important;
297    font-weight: 500 !important;
298    color: var(--text-tertiary) !important;
299    cursor: pointer;
300    transition: color 0.2s ease;
301}
302
303.checkbox-wrap input[type="checkbox"]:checked + label {
304    color: var(--accent) !important;
305}
306
307/* Sliders */
308input[type="range"] {
309    width: 100%;
310    height: 6px;
311    border-radius: 3px;
312    background: linear-gradient(to right, var(--accent), var(--accent-hover));
313    outline: none;
314    -webkit-appearance: none;
315    appearance: none;
316}
317
318input[type="range"]::-webkit-slider-thumb {
319    -webkit-appearance: none;
320    appearance: none;
321    width: 20px;
322    height: 20px;
323    border-radius: 50%;
324    background: var(--accent);
325    cursor: pointer;
326    box-shadow: 0 2px 8px rgba(99, 102, 241, 0.4);
327    transition: all 0.2s ease;
328}
329
330input[type="range"]::-webkit-slider-thumb:hover {
331    width: 24px;
332    height: 24px;
333    box-shadow: 0 4px 12px rgba(99, 102, 241, 0.6);
334}
335
336input[type="range"]::-moz-range-thumb {
337    width: 20px;
338    height: 20px;
339    border-radius: 50%;
340    background: var(--accent);
341    cursor: pointer;
342    border: none;
343    box-shadow: 0 2px 8px rgba(99, 102, 241, 0.4);
344    transition: all 0.2s ease;
345}
346
347input[type="range"]::-moz-range-thumb:hover {
348    width: 24px;
349    height: 24px;
350    box-shadow: 0 4px 12px rgba(99, 102, 241, 0.6);
351}
352
353.range-slider-container {
354    margin: 24px 0;
355}
356
357.slider-labels {
358    display: flex;
359    justify-content: space-between;
360    font-size: 12px;
361    color: var(--text-muted) !important;
362    margin-top: 8px;
363}
364
365.slider-input-box {
366    width: 60px;
367    padding: 8px;
368    background: var(--bg-primary) !important;
369    border: 1px solid var(--border-light) !important;
370    border-radius: 8px !important;
371    color: var(--text-primary) !important;
372    font-size: 14px !important;
373    text-align: right;
374    float: right;
375    margin-top: -35px;
376}
377
378/* Buttons */
379button.primary, button.gr-button-primary {
380    background: linear-gradient(135deg, var(--accent) 0%, var(--accent-hover) 100%) !important;
381    color: var(--text-primary) !important;
382    border: none !important;
383    border-radius: 12px !important;
384    font-size: 16px !important;
385    font-weight: 600 !important;
386    padding: 48px 32px !important;
387    height: auto !important;
388    width: 100% !important;
389    cursor: pointer;
390    transition: all 0.3s ease;
391    box-shadow: 0 4px 6px -1px rgba(99, 102, 241, 0.2);
392}
393
394button.primary:hover, button.gr-button-primary:hover {
395    transform: translateY(-2px);
396    box-shadow: 0 10px 15px -3px rgba(99, 102, 241, 0.3);
397}
398
399button.primary:active, button.gr-button-primary:active {
400    transform: translateY(0);
401}
402
403button.secondary, button.gr-button-secondary {
404    background: var(--bg-secondary) !important;
405    color: var(--text-tertiary) !important;
406    border: 1px solid var(--border-light) !important;
407    border-radius: 12px !important;
408    font-size: 14px !important;
409    font-weight: 500 !important;
410    padding: 12px 24px !important;
411    cursor: pointer;
412    transition: all 0.2s ease;
413}
414
415button.secondary:hover, button.gr-button-secondary:hover {
416    border-color: var(--accent) !important;
417    color: var(--accent) !important;
418    background: rgba(99, 102, 241, 0.05) !important;
419}
420
421/* Text inputs and textareas */
422input[type="text"], textarea, .textbox {
423    background: var(--bg-primary) !important;
424    border: 1px solid var(--border-light) !important;
425    border-radius: 8px !important;
426    color: var(--text-primary) !important;
427    font-family: 'JetBrains Mono', monospace !important;
428    font-size: 13px !important;
429    padding: 12px !important;
430    transition: all 0.2s ease;
431}
432
433input[type="text"]:focus, textarea:focus, .textbox:focus {
434    outline: none;
435    border-color: var(--accent) !important;
436    box-shadow: 0 0 0 3px rgba(99, 102, 241, 0.1);
437}
438
439textarea, .textbox {
440    background: var(--bg-tertiary) !important;
441    color: var(--text-muted) !important;
442    line-height: 1.5;
443}
444
445/* Image containers */
446.output-image-container, .gr-image {
447    background: #000000 !important;
448    border-radius: 8px !important;
449    border: 1px solid var(--border-light) !important;
450    overflow: hidden;
451    aspect-ratio: 16 / 9;
452}
453
454.image-label {
455    font-size: 14px;
456    color: var(--text-muted) !important;
457    margin-bottom: 8px;
458}
459
460.slice-indicator {
461    position: absolute;
462    top: 12px;
463    right: 12px;
464    background: rgba(0, 0, 0, 0.6);
465    color: var(--text-primary);
466    padding: 6px 12px;
467    border-radius: 6px;
468    font-size: 13px;
469    font-family: 'JetBrains Mono', monospace;
470}
471
472/* Processing log */
473.log-container {
474    background: var(--bg-tertiary) !important;
475    border: 1px solid var(--border-light) !important;
476    border-radius: 12px !important;
477    padding: 24px !important;
478    margin-top: 40px;
479}
480
481.log-header {
482    color: var(--accent) !important;
483    font-family: 'Space Grotesk', sans-serif;
484    font-size: 16px;
485    font-weight: 600;
486    margin-bottom: 16px;
487    padding-bottom: 12px;
488    border-bottom: 1px solid var(--border-light);
489}
490
491.log-entry {
492    font-family: 'JetBrains Mono', monospace;
493    font-size: 13px;
494    line-height: 1.6;
495    margin: 8px 0;
496    color: var(--text-muted) !important;
497}
498
499.log-entry.success {
500    color: var(--success) !important;
501}
502
503.log-entry.success::before {
504    content: "✓ ";
505    color: var(--success) !important;
506    font-weight: bold;
507    margin-right: 4px;
508}
509
510/* Row and column spacing */
511.row {
512    gap: 40px !important;
513    margin-bottom: 40px !important;
514}
515
516.column {
517    gap: 32px !important;
518}
519
520/* Focus states for accessibility */
521button:focus-visible, input:focus-visible, textarea:focus-visible {
522    outline: 2px solid var(--accent) !important;
523    outline-offset: 2px;
524}
525
526/* File size text */
527.file-size-text {
528    font-size: 14px;
529    color: var(--text-muted) !important;
530    margin-top: 16px;
531}
532
533/* Status message */
534.status-message {
535    padding: 12px 16px;
536    border-radius: 8px;
537    font-size: 14px;
538    margin-top: 16px;
539}
540
541.status-message.success {
542    background: rgba(16, 185, 129, 0.1);
543    color: var(--success) !important;
544    border: 1px solid rgba(16, 185, 129, 0.3);
545}
546
547.status-message.error {
548    background: rgba(239, 68, 68, 0.1);
549    color: #ef4444 !important;
550    border: 1px solid rgba(239, 68, 68, 0.3);
551}
552
553/* Responsive design */
554@media (max-width: 768px) {
555    .gradio-container {
556        padding: 20px !important;
557    }
558    
559    .row {
560        flex-direction: column !important;
561        gap: 24px !important;
562    }
563    
564    h1 {
565        font-size: 24px;
566    }
567    
568    h2 {
569        font-size: 18px;
570    }
571    
572    .input-container, .output-container {
573        padding: 20px !important;
574    }
575}
576"""
577
578with gr.Blocks(
579    title="MRI Preprocessing Pipeline",
580    css=custom_css,
581    theme=gr.themes.Soft(),
582) as demo:
583    # Main title with emoji
584    gr.HTML("""
585    <div style="margin-bottom: 40px;">
586        <h1 style="display: flex; align-items: center; margin-bottom: 8px;">
587            <span style="margin-right: 8px; font-size: 36px;">🧠</span>
588            MRI Data Preprocessing Pipeline
589        </h1>
590        <p class="info-text" style="margin: 0;">
591            Advanced MRI preprocessing with bias field correction, head mask creation, and interactive visualization
592        </p>
593    </div>
594    """)
595    
596    with gr.Row():
597        # Input section (left column)
598        with gr.Column(scale=1):
599            gr.HTML('<div class="input-container"><h2>Input Section</h2>')
600            
601            uploaded_file = gr.File(
602                label="Upload MRI Image (NIfTI)",
603                file_types=[".nii", ".nii.gz"],
604                file_count="single",
605                type="filepath"
606            )
607            
608            file_info = gr.Textbox(
609                label="File Information",
610                interactive=False,
611                lines=2,
612                value="No file uploaded yet"
613            )
614            
615            gr.HTML('</div>')
616        
617        # Processing options (center-left column)
618        with gr.Column(scale=1):
619            gr.HTML('<div class="input-container"><h2>Processing Options</h2>')
620            
621            intensity_norm_checkbox = gr.Checkbox(
622                value=True,
623                label="Apply Intensity Normalization",
624                show_label=True
625            )
626            
627            mask_creation_checkbox = gr.Checkbox(
628                value=True,
629                label="Apply Head Mask Detection",
630                show_label=True
631            )
632            
633            bias_correction_checkbox = gr.Checkbox(
634                value=True,
635                label="Apply Bias Field Correction",
636                show_label=True
637            )
638            
639            gr.HTML('<div style="margin: 24px 0;"><label style="display: block; margin-bottom: 12px;">Slice to Display (%)</label>')
640            
641            with gr.Row():
642                slice_slider = gr.Slider(
643                    minimum=0,
644                    maximum=100,
645                    value=50,
646                    step=1,
647                    show_label=False,
648                    container=False
649                )
650                slice_value = gr.Textbox(
651                    value="50",
652                    lines=1,
653                    max_lines=1,
654                    interactive=True,
655                    container=False
656                )
657            
658            gr.HTML(
659                '<div class="slider-labels"><span>0%</span><span>100%</span></div></div>'
660            )
661            
662            process_btn = gr.Button(
663                "🔄 Process Image",
664                variant="primary",
665                size="lg",
666                scale=1
667            )
668            
669            gr.HTML('</div>')
670        
671        # Output section (right column - 60% width)
672        with gr.Column(scale=2):
673            gr.HTML('<div class="output-container"><h2>Output Section</h2>')
674            
675            output_image = gr.Image(
676                label="Before/After Comparison",
677                type="pil",
678                interactive=False
679            )
680            
681            processing_log = gr.Textbox(
682                label="Processing Log",
683                lines=10,
684                interactive=False,
685                max_lines=15,
686                elem_classes="log-container"
687            )
688            
689            status_message = gr.Textbox(
690                label="Status",
691                interactive=False,
692                lines=1,
693                max_lines=1
694            )
695            
696            with gr.Row():
697                save_btn = gr.Button(
698                    "💾 Save Processed Image",
699                    variant="secondary",
700                    scale=1
701                )
702                download_file = gr.File(
703                    label="Download Processed Image",
704                    interactive=False,
705                    scale=1
706                )
707            
708            gr.HTML('</div>')
709    # Event handlers
710    def update_file_info(file):
711        """Update file information when file is uploaded"""
712        if file is None:
713            return "No file uploaded yet"
714        
715        try:
716            file_path = file.name
717            file_size = os.path.getsize(file_path)
718            
719            # Format file size
720            if file_size < 1024:
721                size_str = f"{file_size} B"
722            elif file_size < 1024 * 1024:
723                size_str = f"{file_size / 1024:.1f} KB"
724            else:
725                size_str = f"{file_size / (1024 * 1024):.1f} MB"
726            
727            filename = Path(file_path).name
728            return f"📄 {filename}\n{size_str}"
729        except Exception as e:
730            return f"Error reading file: {str(e)}"
731    
732    def update_slice_from_slider(value):
733        """Keep slider and textbox in sync"""
734        return str(int(value))
735    
736    def update_slice_from_textbox(value):
737        """Keep textbox and slider in sync"""
738        try:
739            val = int(value)
740            val = max(0, min(100, val))  # Clamp between 0-100
741            return val
742        except:
743            return 50
744    
745    def process_and_get_slice(file, intensity_norm, mask_create, bias_corr, slice_pct):
746        """Process the MRI image with selected options"""
747        if file is None:
748            return None, "No processing: Please upload an image first", "❌ Error: No file uploaded"
749        
750        try:
751            # Calculate actual slice from percentage
752            raw_img = sitk.ReadImage(file, sitk.sitkFloat32)
753            raw_img_arr = sitk.GetArrayFromImage(raw_img)
754            max_slice = raw_img_arr.shape[0] - 1
755            actual_slice = int((slice_pct / 100.0) * max_slice)
756            
757            img, log, arr, status = preprocess_mri(
758                type('obj', (object,), {'name': file})(),
759                bias_corr, mask_create, intensity_norm, actual_slice
760            )
761            
762            # Store processed array for download
763            demo.processed_array = arr
764            demo.original_path = file
765            
766            return img, log, status
767        except Exception as e:
768            return None, f"Error: {str(e)}", f"❌ Error: {str(e)}"
769    
770    def save_and_download():
771        """Save and download the processed image"""
772        if not hasattr(demo, 'processed_array') or demo.processed_array is None:
773            return None, "No processed image available"
774        
775        output_path, msg = save_processed_image(demo.processed_array, demo.original_path)
776        return output_path, msg
777    
778    # Connect events
779    uploaded_file.change(
780        fn=update_file_info,
781        inputs=[uploaded_file],
782        outputs=[file_info]
783    )
784    
785    slice_slider.change(
786        fn=update_slice_from_slider,
787        inputs=[slice_slider],
788        outputs=[slice_value]
789    )
790    
791    slice_value.change(
792        fn=update_slice_from_textbox,
793        inputs=[slice_value],
794        outputs=[slice_slider]
795    )
796    
797    process_btn.click(
798        fn=process_and_get_slice,
799        inputs=[uploaded_file, intensity_norm_checkbox, mask_creation_checkbox, bias_correction_checkbox, slice_slider],
800        outputs=[output_image, processing_log, status_message]
801    )
802    
803    save_btn.click(
804        fn=save_and_download,
805        outputs=[download_file, status_message]
806    )
807
808
809if __name__ == "__main__":
810    demo.launch(server_name="0.0.0.0", server_port=7860, share=True)
811