JayyyyAWW/MRI_Preprocessing
0
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 