Team Ai
Apppublic

mustehsannisarrao/Masked_Auto_Encoder

sourceHugging Faceupdated 7mo agoView on Hugging Face
1likes
app.py811 linesDownload Raw Back to root
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