CoolFace
Apppublic

cbensimon/screenshot2html

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
app.py260 linesDownload Raw Back to root
1import os2import subprocess3import spaces4import torch5 6import gradio as gr7 8from gradio_client.client import DEFAULT_TEMP_DIR9from playwright.sync_api import sync_playwright10from threading import Thread11from transformers import AutoProcessor, AutoModelForCausalLM, TextIteratorStreamer12from transformers.image_utils import to_numpy_array, PILImageResampling, ChannelDimension13from typing import List14from PIL import Image15 16from transformers.image_transforms import resize, to_channel_dimension_format17 18 19subprocess.run('pip install flash-attn --no-build-isolation', env={'FLASH_ATTENTION_SKIP_CUDA_BUILD': "TRUE"}, shell=True)20 21DEVICE = torch.device("cuda")22PROCESSOR = AutoProcessor.from_pretrained(23    "HuggingFaceM4/VLM_WebSight_finetuned",24)25MODEL = AutoModelForCausalLM.from_pretrained(26    "HuggingFaceM4/VLM_WebSight_finetuned",27    trust_remote_code=True,28    torch_dtype=torch.bfloat16,29).to(DEVICE)30if MODEL.config.use_resampler:31    image_seq_len = MODEL.config.perceiver_config.resampler_n_latents32else:33    image_seq_len = (34        MODEL.config.vision_config.image_size // MODEL.config.vision_config.patch_size35    ) ** 236BOS_TOKEN = PROCESSOR.tokenizer.bos_token37BAD_WORDS_IDS = PROCESSOR.tokenizer(["<image>", "<fake_token_around_image>"], add_special_tokens=False).input_ids38 39 40## Utils41 42def convert_to_rgb(image):43    # `image.convert("RGB")` would only work for .jpg images, as it creates a wrong background44    # for transparent images. The call to `alpha_composite` handles this case45    if image.mode == "RGB":46        return image47 48    image_rgba = image.convert("RGBA")49    background = Image.new("RGBA", image_rgba.size, (255, 255, 255))50    alpha_composite = Image.alpha_composite(background, image_rgba)51    alpha_composite = alpha_composite.convert("RGB")52    return alpha_composite53 54# The processor is the same as the Idefics processor except for the BICUBIC interpolation inside siglip,55# so this is a hack in order to redefine ONLY the transform method56def custom_transform(x):57    x = convert_to_rgb(x)58    x = to_numpy_array(x)59    x = resize(x, (960, 960), resample=PILImageResampling.BILINEAR)60    x = PROCESSOR.image_processor.rescale(x, scale=1 / 255)61    x = PROCESSOR.image_processor.normalize(62        x,63        mean=PROCESSOR.image_processor.image_mean,64        std=PROCESSOR.image_processor.image_std65    )66    x = to_channel_dimension_format(x, ChannelDimension.FIRST)67    x = torch.tensor(x)68    return x69 70## End of Utils71 72 73IMAGE_GALLERY_PATHS = [74    f"example_images/{ex_image}"75    for ex_image in os.listdir(f"example_images")76]77 78 79def install_playwright():80    try:81        subprocess.run(["playwright", "install"], check=True)82        print("Playwright installation successful.")83    except subprocess.CalledProcessError as e:84        print(f"Error during Playwright installation: {e}")85 86install_playwright()87 88 89def add_file_gallery(90    selected_state: gr.SelectData,91    gallery_list: List[str]92):93    return Image.open(gallery_list.root[selected_state.index].image.path)94 95 96def render_webpage(97    html_css_code,98):99    with sync_playwright() as p:100        browser = p.chromium.launch(headless=True)101        context = browser.new_context(102            user_agent=(103                "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/107.0.0.0"104                " Safari/537.36"105            )106        )107        page = context.new_page()108        page.set_content(html_css_code)109        page.wait_for_load_state("networkidle")110        output_path_screenshot = f"{DEFAULT_TEMP_DIR}/{hash(html_css_code)}.png"111        _ = page.screenshot(path=output_path_screenshot, full_page=True)112 113        context.close()114        browser.close()115 116    return Image.open(output_path_screenshot)117 118 119@spaces.GPU(duration=180)120def model_inference(121    image,122):123    if image is None:124        raise ValueError("`image` is None. It should be a PIL image.")125 126    inputs = PROCESSOR.tokenizer(127        f"{BOS_TOKEN}<fake_token_around_image>{'<image>' * image_seq_len}<fake_token_around_image>",128        return_tensors="pt",129        add_special_tokens=False,130    )131    inputs["pixel_values"] = PROCESSOR.image_processor(132        [image],133        transform=custom_transform134    )135    inputs = {136        k: v.to(DEVICE)137        for k, v in inputs.items()138    }139 140    streamer = TextIteratorStreamer(141        PROCESSOR.tokenizer,142        decode_kwargs=dict(143            skip_special_tokens=True144        ),145        skip_prompt=True,146    )147    generation_kwargs = dict(148        inputs,149        bad_words_ids=BAD_WORDS_IDS,150        max_length=4096,151        streamer=streamer,152    )153    thread = Thread(154        target=MODEL.generate,155        kwargs=generation_kwargs,156    )157    thread.start()158    generated_text = ""159    for new_text in streamer:160        generated_text += new_text161        print("before yield")162        # yield generated_text, image163        print("after yield")164 165    # Sanity hack166    generated_text = generated_text.replace("</s>", "")167    rendered_page = render_webpage(generated_text)168    return generated_text, rendered_page169 170generated_html = gr.Code(171    label="Extracted HTML",172    elem_id="generated_html",173)174rendered_html = gr.Image(175    label="Rendered HTML",176    show_download_button=False,177    show_share_button=False,178)179# rendered_html = gr.HTML(180#     label="Rendered HTML"181# )182 183 184css = """185.gradio-container{max-width: 1000px!important}186h1{display: flex;align-items: center;justify-content: center;gap: .25em}187*{transition: width 0.5s ease, flex-grow 0.5s ease}188"""189 190 191with gr.Blocks(title="Screenshot to HTML", theme=gr.themes.Base(), css=css) as demo:192    with gr.Row(equal_height=True):193        with gr.Column(scale=4, min_width=250) as upload_area:194            imagebox = gr.Image(195                type="pil",196                label="Screenshot to extract",197                visible=True,198                sources=["upload", "clipboard"],199            )200            with gr.Group():201                with gr.Row():202                    submit_btn = gr.Button(203                        value="▶️ Submit", visible=True, min_width=120204                    )205                    clear_btn = gr.ClearButton(206                        [imagebox, generated_html, rendered_html], value="🧹 Clear", min_width=120207                    )208                    regenerate_btn = gr.Button(209                        value="🔄 Regenerate", visible=True, min_width=120210                    )211        with gr.Column(scale=4):212            rendered_html.render()213 214    with gr.Row():215        generated_html.render()216 217    with gr.Row():218        template_gallery = gr.Gallery(219            value=IMAGE_GALLERY_PATHS,220            label="Templates Gallery",221            allow_preview=False,222            columns=5,223            elem_id="gallery",224            show_share_button=False,225            height=400,226        )227 228    gr.on(229        triggers=[230            imagebox.upload,231            submit_btn.click,232            regenerate_btn.click,233        ],234        fn=model_inference,235        inputs=[imagebox],236        outputs=[generated_html, rendered_html],237        queue=False,238    )239    regenerate_btn.click(240        fn=model_inference,241        inputs=[imagebox],242        outputs=[generated_html, rendered_html],243        queue=False,244    )245    template_gallery.select(246        fn=add_file_gallery,247        inputs=[template_gallery],248        outputs=[imagebox],249        queue=False,250    ).success(251        fn=model_inference,252        inputs=[imagebox],253        outputs=[generated_html, rendered_html],254        queue=False,255    )256    demo.load(queue=False)257 258demo.queue(max_size=40, api_open=False)259demo.launch(max_threads=400)260