mustehsannisarrao/Masked_Auto_Encoder
1
1import streamlit as st2import torch3import torch.nn as nn4import numpy as np5from PIL import Image6from torchvision import transforms7 8# ── Page Config ──9st.set_page_config(10 page_title="MAE Image Reconstruction",11 page_icon="🎭",12 layout="wide",13 initial_sidebar_state="expanded"14)15 16# ── Custom CSS - Light Mode ──17st.markdown("""18<style>19 /* Import modern fonts */20 @import url('https://fonts.googleapis.com/css2?family=Inter:wght@400;500;600;700;800&display=swap');21 22 /* Global styles */23 * {24 font-family: 'Inter', sans-serif;25 }26 27 /* Main background with subtle gradient */28 .stApp {29 background: linear-gradient(135deg, #f5f7fa 0%, #e9ecf2 100%);30 }31 32 /* Sidebar styling - glass morphism */33 [data-testid="stSidebar"] {34 background: rgba(255, 255, 255, 0.7);35 backdrop-filter: blur(10px);36 border-right: 1px solid rgba(255, 255, 255, 0.3);37 box-shadow: 4px 0 10px rgba(0, 0, 0, 0.02);38 }39 40 [data-testid="stSidebar"] .stMarkdown {41 color: #1a2639;42 }43 44 /* Title styling */45 .main-title {46 text-align: center;47 font-size: 3.2rem;48 font-weight: 800;49 background: linear-gradient(135deg, #667eea 0%, #764ba2 100%);50 -webkit-background-clip: text;51 -webkit-text-fill-color: transparent;52 margin-bottom: 0.2rem;53 letter-spacing: -0.5px;54 text-shadow: 0 2px 10px rgba(102, 126, 234, 0.1);55 }56 57 .sub-title {58 text-align: center;59 color: #64748b;60 font-size: 1rem;61 margin-bottom: 2rem;62 font-weight: 500;63 }64 65 /* Card styling - modern cards */66 .card {67 background: white;68 border: 1px solid rgba(255, 255, 255, 0.8);69 border-radius: 24px;70 padding: 1.5rem;71 box-shadow: 0 10px 30px -5px rgba(0, 0, 0, 0.05), 0 0 0 1px rgba(0, 0, 0, 0.02);72 transition: transform 0.3s ease, box-shadow 0.3s ease;73 backdrop-filter: blur(10px);74 }75 76 .card:hover {77 transform: translateY(-2px);78 box-shadow: 0 20px 40px -10px rgba(102, 126, 234, 0.2), 0 0 0 1px rgba(102, 126, 234, 0.1);79 }80 81 /* Welcome card */82 .welcome-card {83 background: linear-gradient(135deg, #ffffff 0%, #f8fafc 100%);84 border-radius: 32px;85 padding: 4rem 2rem;86 text-align: center;87 box-shadow: 0 25px 50px -12px rgba(0, 0, 0, 0.1);88 border: 1px solid rgba(255, 255, 255, 0.9);89 }90 91 .welcome-emoji {92 font-size: 5rem;93 margin-bottom: 1rem;94 filter: drop-shadow(0 10px 10px rgba(102, 126, 234, 0.2));95 }96 97 .welcome-title {98 font-size: 1.8rem;99 font-weight: 700;100 color: #1e293b;101 margin-bottom: 1rem;102 }103 104 .welcome-text {105 color: #64748b;106 font-size: 1rem;107 max-width: 400px;108 margin: 0 auto;109 }110 111 /* Image labels */112 .img-label {113 text-align: center;114 font-size: 0.9rem;115 font-weight: 600;116 color: #64748b;117 margin-top: 1rem;118 letter-spacing: 0.5px;119 text-transform: uppercase;120 display: flex;121 align-items: center;122 justify-content: center;123 gap: 0.5rem;124 }125 126 .img-label span {127 background: linear-gradient(135deg, #667eea15 0%, #764ba215 100%);128 padding: 0.25rem 0.75rem;129 border-radius: 100px;130 color: #667eea;131 font-size: 0.75rem;132 }133 134 /* Stats box - modern metrics */135 .stat-box {136 background: white;137 border: 1px solid rgba(102, 126, 234, 0.1);138 border-radius: 20px;139 padding: 1.25rem;140 text-align: center;141 box-shadow: 0 4px 6px -1px rgba(0, 0, 0, 0.02);142 }143 144 .stat-number {145 font-size: 2rem;146 font-weight: 800;147 background: linear-gradient(135deg, #667eea 0%, #764ba2 100%);148 -webkit-background-clip: text;149 -webkit-text-fill-color: transparent;150 line-height: 1.2;151 }152 153 .stat-label {154 font-size: 0.75rem;155 color: #94a3b8;156 text-transform: uppercase;157 letter-spacing: 0.5px;158 font-weight: 600;159 }160 161 /* Button styling */162 .stButton > button {163 width: 100%;164 background: linear-gradient(135deg, #667eea 0%, #764ba2 100%);165 color: white;166 border: none;167 border-radius: 16px;168 padding: 0.875rem;169 font-size: 1rem;170 font-weight: 600;171 letter-spacing: 0.5px;172 transition: all 0.3s cubic-bezier(0.4, 0, 0.2, 1);173 box-shadow: 0 4px 6px -1px rgba(102, 126, 234, 0.2), 0 2px 4px -1px rgba(102, 126, 234, 0.1);174 border: 1px solid rgba(255, 255, 255, 0.1);175 }176 177 .stButton > button:hover {178 transform: translateY(-2px);179 box-shadow: 0 20px 25px -5px rgba(102, 126, 234, 0.3), 0 10px 10px -5px rgba(102, 126, 234, 0.1);180 }181 182 .stButton > button:active {183 transform: translateY(0);184 }185 186 /* Slider styling */187 .stSlider {188 padding: 0.5rem 0;189 }190 191 .stSlider > div > div {192 background: linear-gradient(90deg, #667eea, #764ba2) !important;193 }194 195 .stSlider > div > div > div {196 background: white !important;197 border: 2px solid #667eea !important;198 width: 20px !important;199 height: 20px !important;200 box-shadow: 0 2px 4px rgba(0, 0, 0, 0.1) !important;201 }202 203 /* Upload area */204 [data-testid="stFileUploader"] {205 background: white;206 border: 2px dashed #e2e8f0;207 border-radius: 20px;208 padding: 1rem;209 transition: all 0.3s ease;210 }211 212 [data-testid="stFileUploader"]:hover {213 border-color: #667eea;214 background: linear-gradient(135deg, #667eea05 0%, #764ba205 100%);215 }216 217 [data-testid="stFileUploader"] button {218 background: linear-gradient(135deg, #667eea 0%, #764ba2 100%) !important;219 color: white !important;220 border: none !important;221 border-radius: 12px !important;222 padding: 0.5rem 1rem !important;223 font-weight: 500 !important;224 }225 226 /* Info cards */227 .info-card {228 background: white;229 border-radius: 24px;230 padding: 2rem 1.5rem;231 text-align: center;232 height: 100%;233 border: 1px solid rgba(102, 126, 234, 0.1);234 transition: all 0.3s ease;235 }236 237 .info-card:hover {238 border-color: rgba(102, 126, 234, 0.3);239 box-shadow: 0 20px 40px -12px rgba(102, 126, 234, 0.2);240 }241 242 .info-emoji {243 font-size: 2.5rem;244 margin-bottom: 1rem;245 filter: drop-shadow(0 4px 4px rgba(102, 126, 234, 0.2));246 }247 248 .info-title {249 font-weight: 700;250 color: #1e293b;251 margin-bottom: 0.5rem;252 font-size: 1.2rem;253 }254 255 .info-text {256 color: #64748b;257 font-size: 0.9rem;258 line-height: 1.5;259 }260 261 /* Divider */262 .custom-divider {263 background: linear-gradient(90deg, transparent, #667eea, #764ba2, #667eea, transparent);264 height: 2px;265 margin: 2rem 0;266 border: none;267 }268 269 /* Sidebar headers */270 .sidebar-header {271 font-size: 1.2rem;272 font-weight: 700;273 color: #1e293b;274 margin: 1.5rem 0 1rem 0;275 display: flex;276 align-items: center;277 gap: 0.5rem;278 }279 280 .sidebar-header:first-of-type {281 margin-top: 0;282 }283 284 /* Footer */285 .footer {286 color: #94a3b8;287 font-size: 0.75rem;288 text-align: center;289 padding: 2rem 0 1rem 0;290 border-top: 1px solid #e2e8f0;291 margin-top: 2rem;292 }293 294 /* Progress indicator */295 .progress-indicator {296 display: flex;297 align-items: center;298 justify-content: space-between;299 margin: 1rem 0;300 }301 302 .progress-step {303 flex: 1;304 text-align: center;305 position: relative;306 }307 308 .progress-step:not(:last-child):after {309 content: '';310 position: absolute;311 top: 15px;312 right: -50%;313 width: 100%;314 height: 2px;315 background: linear-gradient(90deg, #667eea, #764ba2);316 z-index: 0;317 }318 319 .progress-circle {320 width: 30px;321 height: 30px;322 background: white;323 border: 2px solid #667eea;324 border-radius: 50%;325 margin: 0 auto;326 position: relative;327 z-index: 1;328 display: flex;329 align-items: center;330 justify-content: center;331 font-weight: 600;332 color: #667eea;333 }334 335 .progress-label {336 font-size: 0.7rem;337 color: #64748b;338 margin-top: 0.25rem;339 font-weight: 500;340 }341 342 /* Image container */343 .image-container {344 background: linear-gradient(135deg, #f8fafc 0%, #f1f5f9 100%);345 border-radius: 20px;346 padding: 1rem;347 border: 1px solid #e2e8f0;348 }349 350 /* Hide streamlit branding */351 #MainMenu {visibility: hidden;}352 footer {visibility: hidden;}353 header {visibility: hidden;}354 355 /* Spinner */356 .stSpinner > div {357 border-color: #667eea !important;358 border-top-color: transparent !important;359 }360 361 /* Success message */362 .success-message {363 background: linear-gradient(135deg, #10b98115 0%, #05966915 100%);364 border: 1px solid #10b98130;365 border-radius: 16px;366 padding: 1rem;367 color: #059669;368 font-weight: 500;369 display: flex;370 align-items: center;371 gap: 0.5rem;372 }373</style>374""", unsafe_allow_html=True)375 376# ============================================377# MODEL CLASSES (unchanged)378# ============================================379def get_1d_sincos_pos_embed(embed_dim, pos):380 assert embed_dim % 2 == 0381 omega = np.arange(embed_dim // 2, dtype=np.float32)382 omega /= embed_dim / 2.383 omega = 1. / 10000**omega384 pos = pos.reshape(-1)385 out = np.einsum('m,d->md', pos, omega)386 emb = np.concatenate([np.sin(out), np.cos(out)], axis=1)387 return emb388 389def get_2d_sincos_pos_embed(embed_dim, grid_size):390 assert embed_dim % 2 == 0391 grid_h = np.arange(grid_size, dtype=np.float32)392 grid_w = np.arange(grid_size, dtype=np.float32)393 grid = np.stack(np.meshgrid(grid_w, grid_h), axis=0)394 emb_h = get_1d_sincos_pos_embed(embed_dim // 2, grid[0].flatten())395 emb_w = get_1d_sincos_pos_embed(embed_dim // 2, grid[1].flatten())396 pos_embed = np.concatenate([emb_h, emb_w], axis=1)397 return torch.from_numpy(pos_embed).float()398 399class Patchify(nn.Module):400 def __init__(self, patch_size=16):401 super().__init__()402 self.patch_size = patch_size403 def forward(self, imgs):404 B, C, H, W = imgs.shape405 p = self.patch_size406 imgs = imgs.reshape(B, C, H//p, p, W//p, p)407 imgs = imgs.permute(0, 2, 4, 3, 5, 1).contiguous()408 return imgs.reshape(B, (H//p)*(W//p), p*p*C)409 410class Attention(nn.Module):411 def __init__(self, dim, num_heads=12, qkv_bias=True, dropout=0.1):412 super().__init__()413 self.num_heads = num_heads414 self.head_dim = dim // num_heads415 self.scale = self.head_dim ** -0.5416 self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)417 self.proj = nn.Linear(dim, dim)418 self.dropout = nn.Dropout(dropout)419 self.attn_drop = nn.Dropout(dropout)420 def forward(self, x):421 B, N, C = x.shape422 qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, self.head_dim)423 qkv = qkv.permute(2, 0, 3, 1, 4)424 q, k, v = qkv[0], qkv[1], qkv[2]425 attn = (q @ k.transpose(-2, -1)) * self.scale426 attn = attn.softmax(dim=-1)427 attn = self.attn_drop(attn)428 x = (attn @ v).transpose(1, 2).reshape(B, N, C)429 return self.dropout(self.proj(x))430 431class TransformerBlock(nn.Module):432 def __init__(self, dim, num_heads, mlp_ratio=4.0, dropout=0.1):433 super().__init__()434 self.norm1 = nn.LayerNorm(dim)435 self.attn = Attention(dim, num_heads, dropout=dropout)436 self.norm2 = nn.LayerNorm(dim)437 hidden = int(dim * mlp_ratio)438 self.mlp = nn.Sequential(439 nn.Linear(dim, hidden), nn.GELU(), nn.Dropout(dropout),440 nn.Linear(hidden, dim), nn.Dropout(dropout)441 )442 def forward(self, x):443 x = x + self.attn(self.norm1(x))444 x = x + self.mlp(self.norm2(x))445 return x446 447class RandomMasking(nn.Module):448 def __init__(self, mask_ratio=0.75):449 super().__init__()450 self.mask_ratio = mask_ratio451 def forward(self, x):452 B, N, D = x.shape453 len_keep = int(N * (1 - self.mask_ratio))454 noise = torch.rand(B, N, device=x.device)455 ids_shuffle = torch.argsort(noise, dim=1)456 ids_restore = torch.argsort(ids_shuffle, dim=1)457 ids_keep = ids_shuffle[:, :len_keep]458 x_visible = torch.gather(x, 1, ids_keep.unsqueeze(-1).expand(-1, -1, D))459 mask = torch.ones(B, N, device=x.device)460 mask[:, :len_keep] = 0461 mask = torch.gather(mask, 1, ids_restore)462 return x_visible, mask, ids_restore, ids_keep463 464class MAEEncoder(nn.Module):465 def __init__(self, embed_dim=768, depth=12, num_heads=12, mlp_ratio=4.0, dropout=0.1):466 super().__init__()467 self.blocks = nn.ModuleList([TransformerBlock(embed_dim, num_heads, mlp_ratio, dropout) for _ in range(depth)])468 self.norm = nn.LayerNorm(embed_dim)469 def forward(self, x):470 for blk in self.blocks: x = blk(x)471 return self.norm(x)472 473class MAEDecoder(nn.Module):474 def __init__(self, embed_dim=384, depth=12, num_heads=6, mlp_ratio=4.0,475 dropout=0.1, num_patches=196, patch_size=16):476 super().__init__()477 self.num_patches = num_patches478 self.mask_token = nn.Parameter(torch.zeros(1, 1, embed_dim))479 nn.init.trunc_normal_(self.mask_token, std=0.02)480 grid_size = int(num_patches ** 0.5)481 dec_pos = get_2d_sincos_pos_embed(embed_dim, grid_size)482 self.register_buffer('pos_embed', dec_pos.unsqueeze(0))483 self.blocks = nn.ModuleList([TransformerBlock(embed_dim, num_heads, mlp_ratio, dropout) for _ in range(depth)])484 self.norm = nn.LayerNorm(embed_dim)485 self.pred = nn.Linear(embed_dim, patch_size * patch_size * 3)486 def forward(self, x, ids_restore):487 B, len_keep, _ = x.shape488 mask_tokens = self.mask_token.repeat(B, self.num_patches - len_keep, 1)489 x = torch.cat([x, mask_tokens], dim=1)490 x = torch.gather(x, 1, ids_restore.unsqueeze(-1).expand(-1, -1, x.shape[-1]))491 x = x + self.pos_embed492 for blk in self.blocks: x = blk(x)493 return self.pred(self.norm(x))494 495class MaskedAutoencoder(nn.Module):496 def __init__(self, img_size=224, patch_size=16, in_chans=3,497 encoder_embed_dim=768, encoder_depth=12, encoder_num_heads=12,498 decoder_embed_dim=384, decoder_depth=12, decoder_num_heads=6,499 mask_ratio=0.75, mlp_ratio=4.0, dropout=0.1):500 super().__init__()501 self.patch_size = patch_size502 self.num_patches = (img_size // patch_size) ** 2503 self.patchify = Patchify(patch_size)504 self.patch_embed = nn.Linear(patch_size * patch_size * in_chans, encoder_embed_dim)505 pos = get_2d_sincos_pos_embed(encoder_embed_dim, img_size // patch_size)506 self.register_buffer('pos_embed', pos.unsqueeze(0))507 self.masking = RandomMasking(mask_ratio)508 self.encoder = MAEEncoder(encoder_embed_dim, encoder_depth, encoder_num_heads, mlp_ratio, dropout)509 self.enc_to_dec = nn.Linear(encoder_embed_dim, decoder_embed_dim)510 self.decoder = MAEDecoder(decoder_embed_dim, decoder_depth, decoder_num_heads,511 mlp_ratio, dropout, self.num_patches, patch_size)512 nn.init.xavier_uniform_(self.patch_embed.weight)513 nn.init.constant_(self.patch_embed.bias, 0)514 nn.init.xavier_uniform_(self.enc_to_dec.weight)515 nn.init.constant_(self.enc_to_dec.bias, 0)516 517 def forward(self, imgs):518 x = self.patch_embed(self.patchify(imgs))519 x = x + self.pos_embed520 x_visible, mask, ids_restore, ids_keep = self.masking(x)521 latent = self.encoder(x_visible)522 latent = self.enc_to_dec(latent)523 pred = self.decoder(latent, ids_restore)524 return pred, mask, ids_restore525 526 def reconstruct_image(self, pred):527 B, N, _ = pred.shape528 p = self.patch_size529 h = w = int(N ** 0.5)530 pred = pred.reshape(B, h, w, p, p, 3)531 pred = pred.permute(0, 5, 1, 3, 2, 4).contiguous()532 return pred.reshape(B, 3, h*p, w*p)533 534# ============================================535# LOAD MODEL536# ============================================537DEVICE = torch.device('cpu')538 539@st.cache_resource540def load_model():541 model = MaskedAutoencoder().to(DEVICE)542 model.load_state_dict(torch.load('best_model.pth', map_location='cpu'))543 model.eval()544 return model545 546# ============================================547# RECONSTRUCT548# ============================================549def reconstruct(image, mask_ratio, model):550 transform = transforms.Compose([551 transforms.Resize(224),552 transforms.CenterCrop(224),553 transforms.ToTensor(),554 transforms.Normalize(mean=[0.485, 0.456, 0.406],555 std=[0.229, 0.224, 0.225])556 ])557 img_tensor = transform(image).unsqueeze(0).to(DEVICE)558 model.masking.mask_ratio = mask_ratio559 MEAN = np.array([0.485, 0.456, 0.406])560 STD = np.array([0.229, 0.224, 0.225])561 with torch.no_grad():562 pred, mask, _ = model(img_tensor)563 patches_raw = model.patchify(img_tensor)564 masked_patches = patches_raw.clone()565 masked_patches[mask.bool()] = 0566 masked_img = model.reconstruct_image(masked_patches)567 recon_img = model.reconstruct_image(pred)568 def to_pil(t):569 img_np = t[0].cpu().numpy().transpose(1, 2, 0)570 img_np = img_np * STD + MEAN571 img_np = np.clip(img_np, 0, 1)572 return Image.fromarray((img_np * 255).astype(np.uint8))573 return to_pil(masked_img), to_pil(recon_img)574 575# ============================================576# UI577# ============================================578 579# Header with animation effect580st.markdown("""581<div style='text-align: center; animation: fadeIn 1s ease-in;'>582 <div class="main-title">🎭 MAE Image Reconstruction</div>583 <div class="sub-title">Masked Autoencoder · ViT-Base/16 · TinyImageNet</div>584</div>585 586<style>587@keyframes fadeIn {588 from { opacity: 0; transform: translateY(-20px); }589 to { opacity: 1; transform: translateY(0); }590}591</style>592""", unsafe_allow_html=True)593 594# Load model595model = load_model()596 597# Sidebar598with st.sidebar:599 st.markdown("""600 <div class="sidebar-header">601 <span>⚙️</span> Controls602 </div>603 """, unsafe_allow_html=True)604 605 st.markdown("<hr class='custom-divider'>", unsafe_allow_html=True)606 607 uploaded = st.file_uploader(608 "📁 Upload Image",609 type=['jpg', 'jpeg', 'png'],610 help="JPG, JPEG, PNG supported (Max 200MB)"611 )612 613 st.markdown("<hr class='custom-divider'>", unsafe_allow_html=True)614 615 st.markdown("""616 <div class="sidebar-header">617 <span>🎛️</span> Masking Ratio618 </div>619 """, unsafe_allow_html=True)620 621 mask_ratio = st.slider(622 "", min_value=0.1, max_value=0.9,623 value=0.75, step=0.05,624 help="Higher = more patches masked"625 )626 627 masked_pct = int(mask_ratio * 100)628 visible_pct = 100 - masked_pct629 630 # Progress indicator631 st.markdown("""632 <div class="progress-indicator">633 <div class="progress-step">634 <div class="progress-circle">🖼️</div>635 <div class="progress-label">Original</div>636 </div>637 <div class="progress-step">638 <div class="progress-circle">🎭</div>639 <div class="progress-label">Masked</div>640 </div>641 <div class="progress-step">642 <div class="progress-circle">✨</div>643 <div class="progress-label">Reconstructed</div>644 </div>645 </div>646 """, unsafe_allow_html=True)647 648 col1, col2 = st.columns(2)649 with col1:650 st.markdown(f"""651 <div class="stat-box">652 <div class="stat-number">{masked_pct}%</div>653 <div class="stat-label">Masked</div>654 </div>655 """, unsafe_allow_html=True)656 with col2:657 st.markdown(f"""658 <div class="stat-box">659 <div class="stat-number">{visible_pct}%</div>660 <div class="stat-label">Visible</div>661 </div>662 """, unsafe_allow_html=True)663 664 st.markdown("<hr class='custom-divider'>", unsafe_allow_html=True)665 666 run_btn = st.button("🚀 Reconstruct Image", use_container_width=True)667 668 st.markdown("""669 <div class="footer">670 MAE · He et al. 2021<br>671 ViT-Base Encoder · ViT-Small Decoder<br>672 Trained on TinyImageNet673 </div>674 """, unsafe_allow_html=True)675 676# Main content677if uploaded is None:678 # Welcome screen with modern design679 st.markdown("""680 <div class="welcome-card">681 <div class="welcome-emoji">🖼️</div>682 <div class="welcome-title">Ready to Transform Your Images?</div>683 <div class="welcome-text">684 Upload an image and watch as our MAE model magically reconstructs 685 it from randomly masked patches using advanced Vision Transformer technology.686 </div>687 </div>688 """, unsafe_allow_html=True)689 690 # Info cards in a grid691 st.markdown("<br>", unsafe_allow_html=True)692 693 col1, col2, col3 = st.columns(3)694 695 with col1:696 st.markdown("""697 <div class="info-card">698 <div class="info-emoji">🎭</div>699 <div class="info-title">Random Masking</div>700 <div class="info-text">Up to 90% of image patches are randomly masked, challenging the model to understand context</div>701 </div>702 """, unsafe_allow_html=True)703 704 with col2:705 st.markdown("""706 <div class="info-card">707 <div class="info-emoji">🧠</div>708 <div class="info-title">Vision Transformer</div>709 <div class="info-text">12-layer ViT-Base encoder with 768-dim embeddings for powerful feature extraction</div>710 </div>711 """, unsafe_allow_html=True)712 713 with col3:714 st.markdown("""715 <div class="info-card">716 <div class="info-emoji">⚡</div>717 <div class="info-title">Real-time Processing</div>718 <div class="info-text">Instant reconstruction with adjustable masking ratio for interactive experimentation</div>719 </div>720 """, unsafe_allow_html=True)721 722else:723 image = Image.open(uploaded).convert('RGB')724 725 if run_btn:726 with st.spinner('🔄 Reconstructing your image...'):727 masked_img, recon_img = reconstruct(image, mask_ratio, model)728 729 st.markdown("""730 <div class="success-message">731 <span>✨</span> Reconstruction complete! Compare the results below.732 </div>733 """, unsafe_allow_html=True)734 735 st.markdown("<br>", unsafe_allow_html=True)736 737 # Create three columns with custom styling738 col1, col2, col3 = st.columns(3)739 740 with col1:741 st.markdown('<div class="card">', unsafe_allow_html=True)742 st.markdown('<div class="image-container">', unsafe_allow_html=True)743 st.image(image, use_column_width=True)744 st.markdown('</div>', unsafe_allow_html=True)745 st.markdown("""746 <div class="img-label">747 <span>📷</span> Original Image748 </div>749 """, unsafe_allow_html=True)750 st.markdown('</div>', unsafe_allow_html=True)751 752 with col2:753 st.markdown('<div class="card">', unsafe_allow_html=True)754 st.markdown('<div class="image-container">', unsafe_allow_html=True)755 st.image(masked_img, use_column_width=True)756 st.markdown('</div>', unsafe_allow_html=True)757 st.markdown(f"""758 <div class="img-label">759 <span>🎭</span> Masked ({masked_pct}% hidden)760 </div>761 """, unsafe_allow_html=True)762 st.markdown('</div>', unsafe_allow_html=True)763 764 with col3:765 st.markdown('<div class="card">', unsafe_allow_html=True)766 st.markdown('<div class="image-container">', unsafe_allow_html=True)767 st.image(recon_img, use_column_width=True)768 st.markdown('</div>', unsafe_allow_html=True)769 st.markdown("""770 <div class="img-label">771 <span>✨</span> Reconstructed772 </div>773 """, unsafe_allow_html=True)774 st.markdown('</div>', unsafe_allow_html=True)775 776 # Add download buttons777 st.markdown("<br>", unsafe_allow_html=True)778 col1, col2, col3 = st.columns(3)779 780 with col2:781 st.download_button(782 label="📥 Download Reconstructed Image",783 data=recon_img.tobytes(),784 file_name="reconstructed.png",785 mime="image/png",786 use_container_width=True787 )788 789 else:790 st.markdown("<br>", unsafe_allow_html=True)791 st.markdown("""792 <div style='display: flex; justify-content: center;'>793 <div class="card" style='max-width: 500px; width: 100%;'>794 <div class="image-container" style='text-align: center;'>795 """, unsafe_allow_html=True)796 797 st.image(image, width=400)798 799 st.markdown("""800 </div>801 <div class="img-label" style='margin-top: 1rem;'>802 <span>📷</span> Preview803 </div>804 <p style='text-align: center; color: #64748b; margin-top: 0.5rem;'>805 Click "Reconstruct Image" to see the magic happen!806 </p>807 </div>808 </div>809 """, unsafe_allow_html=True)810 811 #hi