File size: 4,673 Bytes
1cca2dd
1a9eb8a
 
 
 
 
 
1cca2dd
 
1a9eb8a
 
1cca2dd
 
1a9eb8a
 
 
 
1cca2dd
 
 
 
 
 
 
 
 
 
 
 
 
1a9eb8a
1cca2dd
 
 
 
1a9eb8a
1cca2dd
 
 
 
 
1a9eb8a
1cca2dd
 
1a9eb8a
1cca2dd
 
 
 
1a9eb8a
1cca2dd
 
1a9eb8a
1cca2dd
 
 
 
 
 
1a9eb8a
1cca2dd
 
 
1a9eb8a
1cca2dd
 
1a9eb8a
1cca2dd
 
 
 
 
1a9eb8a
1cca2dd
 
 
 
 
 
1a9eb8a
f0a94d4
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
"""Pinned-release Gradio app. Import spaces before torch for ZeroGPU emulation."""
import os
from pathlib import Path
import spaces
import torch
import gradio as gr
from huggingface_hub import snapshot_download
from inference import prepare, chroma_coefficients, render, load_colorizer
import json

REPO=os.getenv('MODEL_ID','User-2468/mini-unet-colorizer')
REV=os.getenv('MODEL_REVISION','1a9eb8af2754ad2329a24cfe50d388cb559441d0')
SUBFOLDER=os.getenv('MODEL_SUBFOLDER','').strip('/')
CPU=os.getenv('ZEROGPU_CPU_TEST')=='1'
DEVICE='cpu' if CPU else 'cuda'
if CPU:torch.set_num_threads(min(4,os.cpu_count() or 1))
path=Path(REPO)
if not path.is_dir():
    prefix=SUBFOLDER+'/' if SUBFOLDER else ''
    path=Path(snapshot_download(REPO,revision=REV,allow_patterns=[prefix+'model.safetensors',prefix+'config.json']))
if SUBFOLDER:path=path/SUBFOLDER
cfg=json.loads((path/'config.json').read_text())
if cfg.get('architecture')=='DecisionColorizer':
    from decision_model import load_decision
    MODEL=load_decision(path,DEVICE)
else:MODEL=load_colorizer(path,DEVICE)
MODEL.eval().requires_grad_(False)
PARAMETERS=sum(p.numel() for p in MODEL.parameters())
if PARAMETERS>=4_000_000:raise RuntimeError('Model exceeds the four-million-parameter limit.')
MODES=getattr(MODEL,'modes',1)

class PaletteView(torch.nn.Module):
    def __init__(self,model,mode):super().__init__();self.model=model;self.mode=mode
    def forward(self,L):
        return self.model(L) if self.mode<0 else self.model(L,True)['all'][:,self.mode]

@spaces.GPU(duration=15)
def infer(small,radius,mode):
    return chroma_coefficients(PaletteView(MODEL,mode),small,radius)

def run(image,strength,smoothing,detail,mode):
    if image is None:raise gr.Error('Upload a photograph first.')
    mode=int(mode)
    if mode < -1 or mode>=MODES or (MODES==1 and mode!=-1):raise gr.Error('That colour interpretation is unavailable.')
    try:
        radius={'Gentle':4,'Balanced':8,'Strong':16}[smoothing]
        size={'Standard':256,'Detailed':384,'Maximum':512}[detail]
        light,alpha,small=prepare(image,size)
        coefficients=infer(small,radius,mode)
        result=render(light,alpha,coefficients,strength)
        cache=(light,alpha,coefficients)
    except (ValueError,RuntimeError,TypeError,KeyError) as exc:
        raise gr.Error(str(exc)) from exc
    return result,cache

def recolour(strength,cache):
    if cache is None:return gr.skip()
    light,alpha,coefficients=cache
    return render(light,alpha,coefficients,strength)

with gr.Blocks(title='Mini Photo Colorizer',delete_cache=(3600,3600)) as demo:
    gr.Markdown('# Give an old photograph new colour\nPlausible colour, with the original detail and proportions preserved.')
    cache=gr.State(value=None,time_to_live=600)
    with gr.Row():
        source=gr.Image(type='pil',image_mode='RGBA',label='Your photograph',sources=['upload','clipboard'],height=460)
        result=gr.Image(type='pil',label='Colourised photograph',format='png',height=460,interactive=False)
    with gr.Row():
        button=gr.Button('Colourise',variant='primary',scale=2)
        clear=gr.ClearButton([source,result,cache],value='Clear',scale=1)
    strength=gr.Slider(0,1.5,value=1.,step=.05,label='Colour strength',info='Adjust after colourising without another GPU request.')
    with gr.Accordion('Fine-tune the result',open=False):
        detail=gr.Radio(['Standard','Detailed','Maximum'],value='Standard',label='Detail',info='Higher detail can change the colours. Try Detailed if small objects are missed.')
        smoothing=gr.Radio(['Gentle','Balanced','Strong'],value='Balanced',label='Colour smoothing')
        palette=gr.Dropdown([('Automatic',-1)]+[(f'Alternative {i+1}',i) for i in range(MODES)] if MODES>1 else [('Automatic',-1)],value=-1,label='Colour interpretation',visible=MODES>1)
        gr.Markdown('After changing these settings, select **Colourise** again.')
    button.click(run,[source,strength,smoothing,detail,palette],[result,cache],api_name='colorize',concurrency_limit=1)
    strength.release(recolour,[strength,cache],result,api_name=False,concurrency_limit=1)
    source.change(lambda:None,outputs=cache,api_name=False,queue=False)
    gr.Markdown('Download the result as PNG using its download button. Up to 12 megapixels; transparency is preserved.\n\nColours are interpretations, not recovered historical facts. Unusual scenes, damaged scans and tiny objects can still produce muted colours or colour bleeding. [About the model](https://huggingface.co/User-2468/mini-unet-colorizer).')
demo.queue(default_concurrency_limit=1,max_size=12)
if __name__=='__main__':demo.launch(max_file_size='30mb',ssr_mode=False)