""" Deep Analog — Inference Pipeline (deployment build) Privacy by design: every function in this module operates purely on in-memory bytes and tensors. Nothing is ever written to disk, no filenames or image contents are logged, and all buffers are released when the request completes. """ import io import os import torch import torch.nn.functional as F from PIL import Image import torchvision.transforms as transforms # Optional: RAW and HEIF support try: import rawpy HAS_RAWPY = True except ImportError: HAS_RAWPY = False try: import pillow_heif pillow_heif.register_heif_opener() HAS_HEIF = True except ImportError: HAS_HEIF = False from models import ( StyleLUTNet, match_tone_curve, render_grain, render_halation, ) # ========================================================================= # Format detection helpers # ========================================================================= RAW_EXTENSIONS = { '.raf', '.cr2', '.cr3', '.nef', '.nrw', '.arw', '.srf', '.sr2', '.dng', '.orf', '.rw2', '.pef', '.srw', '.x3f', '.erf', '.mrw', '.3fr', '.mos', '.mef', '.iiq', '.rwl', '.kdc', '.dcr', } HEIF_EXTENSIONS = {'.heif', '.heic', '.hif', '.avif'} def _detect_format_from_bytes(file_bytes: bytes) -> str: if len(file_bytes) < 12: return 'standard' if file_bytes[4:8] == b'ftyp': brand = file_bytes[8:12] heif_brands = [b'heic', b'heix', b'hevc', b'hevx', b'mif1', b'msf1', b'avif'] if brand in heif_brands: return 'heif' if file_bytes[:2] in (b'II', b'MM'): if len(file_bytes) > 10 and file_bytes[8:10] == b'CR': return 'raw' return 'raw_or_tiff' if file_bytes[:16].startswith(b'FUJIFILMCCD-RAW'): return 'raw' if file_bytes[:4] in (b'IIRO', b'MMOR'): return 'raw' if file_bytes[:4] == b'IIU\x00': return 'raw' return 'standard' def _open_raw(file_bytes: bytes) -> Image.Image: if not HAS_RAWPY: raise RuntimeError('RAW support is not available on this server.') raw = rawpy.imread(io.BytesIO(file_bytes)) rgb = raw.postprocess( use_camera_wb=True, half_size=False, no_auto_bright=False, output_bps=8, ) return Image.fromarray(rgb) def _open_image_any(file_bytes: bytes, filename: str = '') -> Image.Image: """Open image from bytes. Supports JPEG/PNG/BMP/WebP/GIF/TIFF, RAW (RAF, CR2, CR3, NEF, ARW, DNG, ...), and HEIF/HEIC/AVIF.""" ext = os.path.splitext(filename)[1].lower() if filename else '' if ext in RAW_EXTENSIONS: return _open_raw(file_bytes) if ext in HEIF_EXTENSIONS: if not HAS_HEIF: raise RuntimeError('HEIF/HEIC support is not available on this server.') return Image.open(io.BytesIO(file_bytes)).convert('RGB') fmt = _detect_format_from_bytes(file_bytes) if fmt == 'raw': return _open_raw(file_bytes) if fmt == 'raw_or_tiff': if HAS_RAWPY: try: return _open_raw(file_bytes) except Exception: pass return Image.open(io.BytesIO(file_bytes)).convert('RGB') if fmt == 'heif': if not HAS_HEIF: raise RuntimeError('HEIF/HEIC support is not available on this server.') return Image.open(io.BytesIO(file_bytes)).convert('RGB') return Image.open(io.BytesIO(file_bytes)).convert('RGB') # ========================================================================= # Reference tone analysis + film tone rendering # ========================================================================= def analyze_reference_tone(ref_tensor): img = ref_tensor.squeeze(0) lum = 0.2126 * img[0] + 0.7152 * img[1] + 0.0722 * img[2] dark_thresh = torch.quantile(lum, 0.02).item() bright_thresh = torch.quantile(lum, 0.98).item() contrast_range = bright_thresh - dark_thresh shadow_mask = (lum < torch.quantile(lum, 0.15)).float() shadow_pixels = shadow_mask.sum().clamp(min=1) shadow_color = torch.stack([ (img[c] * shadow_mask).sum() / shadow_pixels for c in range(3) ]) highlight_mask = (lum > torch.quantile(lum, 0.85)).float() highlight_pixels = highlight_mask.sum().clamp(min=1) highlight_color = torch.stack([ (img[c] * highlight_mask).sum() / highlight_pixels for c in range(3) ]) gray = img.mean(dim=0, keepdim=True) chroma = (img - gray).pow(2).sum(dim=0).sqrt().mean().item() return { 'black_point': dark_thresh, 'white_point': bright_thresh, 'contrast_range': contrast_range, 'shadow_color': shadow_color, 'highlight_color': highlight_color, 'saturation': chroma, } def apply_film_tone(img, ref_tone, strength=0.7): B, C, H, W = img.shape dev = img.device ref_black = ref_tone['black_point'] if ref_black > 0.02: black_lift = ref_black * strength * 0.8 img = img * (1.0 - black_lift) + black_lift ref_white = ref_tone['white_point'] if ref_white < 0.95: compression = 1.0 - (1.0 - ref_white) * strength knee = 0.6 above_knee = ((img - knee) / (1.0 - knee)).clamp(0, 1) rolloff_amount = above_knee.pow(2) * (3.0 - 2.0 * above_knee) img = img * (1.0 - rolloff_amount * (1.0 - compression) * 0.6) lum = (0.2126 * img[:, 0:1] + 0.7152 * img[:, 1:2] + 0.0722 * img[:, 2:3]) shadow_weight = (1.0 - lum).clamp(0, 1).pow(2) ref_shadow = ref_tone['shadow_color'].to(dev) shadow_mean = ref_shadow.mean() shadow_tint = (ref_shadow - shadow_mean).view(1, 3, 1, 1) img = img + shadow_tint * shadow_weight * strength * 0.5 highlight_weight = lum.clamp(0, 1).pow(2) ref_highlight = ref_tone['highlight_color'].to(dev) highlight_mean = ref_highlight.mean() highlight_tint = (ref_highlight - highlight_mean).view(1, 3, 1, 1) img = img + highlight_tint * highlight_weight * strength * 0.3 dither = (torch.rand_like(img) + torch.rand_like(img) - 1.0) * (0.4 / 256.0) img = img + dither return img.clamp(0, 1) # ========================================================================= # Pipeline # ========================================================================= # Default film parameters (typical 35mm color negative characteristics) DEFAULT_FILM_PARAMS = { 'sigma': 0.018, 'grain_size': 1.6, 'lum_a': -2.0, 'lum_b': 1.5, 'lum_c': 0.8, 'h_threshold': 0.72, 'h_radius': 4.5, 'h_intensity': 0.15, 'h_color_bias': [0.3, -0.1, -0.3], } class DeepAnalogPipeline: """Stateful pipeline: loads the model once, runs inference on demand.""" def __init__(self, style_lut_path: str, device: str = 'auto', max_side: int = 2048): if device == 'auto': self.device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') else: self.device = torch.device(device) self.max_side = max_side print(f'[Pipeline] Device: {self.device}') print('[Pipeline] Loading StyleLUT...') self.style_lut = StyleLUTNet(lut_dim=33, lut_dim_low=17) ckpt = torch.load(style_lut_path, map_location='cpu', weights_only=True) state = ckpt['model_state_dict'] if 'model_state_dict' in ckpt else ckpt # Slim checkpoints may be stored in fp16 — cast back to fp32 state = {k: v.float() if torch.is_floating_point(v) else v for k, v in state.items()} self.style_lut.load_state_dict(state) self.style_lut.to(self.device).eval() print('[Pipeline] StyleLUT ready') # ------------------------------------------------------------------ # Helpers # ------------------------------------------------------------------ def _load_image(self, file_bytes: bytes, filename: str = '', max_side: int = None) -> torch.Tensor: img = _open_image_any(file_bytes, filename) cap = max_side if max_side is not None else self.max_side w, h = img.size if cap and max(w, h) > cap: s = cap / max(w, h) img = img.resize((int(w * s), int(h * s)), Image.LANCZOS) return transforms.ToTensor()(img).unsqueeze(0).to(self.device) def _tensor_to_jpeg_bytes(self, tensor: torch.Tensor, quality: int = 95) -> bytes: img = tensor.squeeze(0).cpu().clamp(0, 1) pil = transforms.ToPILImage()(img) buf = io.BytesIO() # Re-encoding strips all metadata (EXIF, GPS, camera serials) pil.save(buf, format='JPEG', quality=quality, subsampling=0) return buf.getvalue() def _lut_to_cube(self, lut: torch.Tensor, title: str = 'Deep Analog StyleLUT') -> str: lut = lut.squeeze(0).cpu().clamp(0, 1) D = lut.shape[1] lines = [ f'TITLE "{title}"', f'LUT_3D_SIZE {D}', 'DOMAIN_MIN 0.0 0.0 0.0', 'DOMAIN_MAX 1.0 1.0 1.0', '', ] for b_idx in range(D): for g_idx in range(D): for r_idx in range(D): r = lut[0, r_idx, g_idx, b_idx].item() g = lut[1, r_idx, g_idx, b_idx].item() b = lut[2, r_idx, g_idx, b_idx].item() lines.append(f'{r:.6f} {g:.6f} {b:.6f}') return '\n'.join(lines) def _make_test_chart(self, width=768, height=128) -> torch.Tensor: """Neutral test chart: hue sweep over the top, gray ramp below. Used to preview what a predicted LUT does, without any user image.""" import colorsys rows = [] hue_h = height * 3 // 4 x = torch.linspace(0, 1, width) # Hue bands at three luminance levels for v in (0.85, 0.6, 0.35): band = torch.zeros(3, hue_h // 3, width) for i in range(width): r, g, b = colorsys.hsv_to_rgb(x[i].item(), 0.75, v) band[0, :, i] = r band[1, :, i] = g band[2, :, i] = b rows.append(band) # Gray ramp ramp = x.view(1, 1, width).expand(3, height - 3 * (hue_h // 3), width).clone() rows.append(ramp) chart = torch.cat(rows, dim=1) return chart.unsqueeze(0).to(self.device) def _film_params_dict(self, fp, ref_tone): def _val(x): return x.item() if isinstance(x, torch.Tensor) else x def _color_list(x): if isinstance(x, torch.Tensor): x = x.detach().cpu().flatten().tolist() return [round(float(v), 4) for v in x] return { 'grain_sigma': round(_val(fp['sigma']), 6), 'grain_size': round(_val(fp['grain_size']), 4), 'lum_a': round(_val(fp['lum_a']), 4), 'lum_b': round(_val(fp['lum_b']), 4), 'lum_c': round(_val(fp['lum_c']), 4), 'h_threshold': round(_val(fp['h_threshold']), 4), 'h_radius': round(_val(fp['h_radius']), 4), 'h_intensity': round(_val(fp['h_intensity']), 4), 'h_color_bias': _color_list(fp['h_color_bias']), 'residual_scale': round(self.style_lut.residual_scale.item(), 4), 'ref_black_point': round(ref_tone['black_point'], 4), 'ref_white_point': round(ref_tone['white_point'], 4), 'ref_contrast_range': round(ref_tone['contrast_range'], 4), 'ref_shadow_color': _color_list(ref_tone['shadow_color']), 'ref_highlight_color': _color_list(ref_tone['highlight_color']), } # ------------------------------------------------------------------ # Editor-preset exports (.xmp for Adobe, .costyle for Capture One) # Both formats hold 1D per-channel curves, not a 3D LUT, so they are # faithful for tone/tint on neutrals; exact color lives in the .cube. # ------------------------------------------------------------------ @torch.no_grad() def _neutral_curves(self, lut, tone_transfer=None, tone_strength=0.7, ref_tone=None, film_tone_strength=0.7, n=16): """Push a neutral gray ramp through the color stages and sample the resulting per-channel response as (x, y) points in 0-255 space.""" g = torch.linspace(0, 1, 256, device=self.device) ramp = g.view(1, 1, 1, 256).expand(1, 3, 1, 256).contiguous() out = self.style_lut.trilinear(lut, ramp).clamp(0, 1) if tone_transfer is not None: idx = out * 255.0 lo = idx.floor().long().clamp(0, 254) frac = idx - lo.float() for c in range(3): tt = tone_transfer[c] mapped = tt[lo[0, c, 0]] * (1 - frac[0, c, 0]) + \ tt[lo[0, c, 0] + 1] * frac[0, c, 0] out[0, c, 0] = out[0, c, 0] * (1 - tone_strength) + \ mapped * tone_strength out = out.clamp(0, 1) if ref_tone is not None: out = apply_film_tone(out, ref_tone, strength=film_tone_strength) xs = [round(i * 255 / (n - 1)) for i in range(n)] curves = {} for c, name in enumerate('rgb'): ch = out[0, c, 0].cpu() curves[name] = [(x, int(round(float(ch[min(x, 255)]) * 255))) for x in xs] return curves @staticmethod def _grain_settings(fp, grain_mult=1.0): amount = int(min(100, round(fp['sigma'] * grain_mult / 0.018 * 25))) size = int(min(100, round(fp['grain_size'] / 1.6 * 25))) return amount, size @staticmethod def _curves_to_xmp(curves, grain_amount, grain_size, name='Deep Analog Look'): import uuid def seq(pts): return ''.join(f' {x}, {y}\n' for x, y in pts) master = [(0, 0), (255, 255)] return f''' {seq(master)} {seq(curves['r'])} {seq(curves['g'])} {seq(curves['b'])} ''' @staticmethod def _curves_to_costyle(curves, grain_amount, grain_size, name='Deep Analog Look'): def pts(ps): return ';'.join(f'{x},{y}' for x, y in ps) c1_grain = int(min(1000, grain_amount * 10)) c1_gran = int(min(1000, grain_size * 10)) return f''' \t \t \t \t \t \t \t \t ''' # ------------------------------------------------------------------ # Mode 1: reference only → predicted LUT # ------------------------------------------------------------------ @torch.no_grad() def process_reference_only(self, reference_bytes: bytes, reference_filename: str = '') -> dict: reference = self._load_image(reference_bytes, filename=reference_filename) ref_224 = F.interpolate(reference, size=(224, 224), mode='bilinear', align_corners=False) lut, _ = self.style_lut.predict_lut(ref_224) cube_text = self._lut_to_cube(lut) ref_tone = analyze_reference_tone(reference) params = self._film_params_dict(DEFAULT_FILM_PARAMS, ref_tone) # LUT preview: identity chart vs LUT-applied chart chart = self._make_test_chart() chart_graded = self.style_lut.trilinear(lut, chart) curves = self._neutral_curves(lut, ref_tone=ref_tone, film_tone_strength=0.7) ga, gs = self._grain_settings(DEFAULT_FILM_PARAMS) return { 'cube_lut': cube_text, 'params': params, 'xmp': self._curves_to_xmp(curves, ga, gs), 'costyle': self._curves_to_costyle(curves, ga, gs), 'chart_before_jpg': self._tensor_to_jpeg_bytes(chart, quality=90), 'chart_after_jpg': self._tensor_to_jpeg_bytes(chart_graded, quality=90), } # ------------------------------------------------------------------ # Mode 2: reference + target → LUT + fully rendered image # ------------------------------------------------------------------ @torch.no_grad() def process( self, reference_bytes: bytes, target_bytes: bytes, tone_strength: float = 0.7, grain_mult: float = 1.0, film_tone_strength: float = 0.7, halation_mult: float = 1.0, reference_filename: str = '', target_filename: str = '', ) -> dict: reference = self._load_image(reference_bytes, filename=reference_filename) target = self._load_image(target_bytes, filename=target_filename) ref_224 = F.interpolate(reference, size=(224, 224), mode='bilinear', align_corners=False) _, _, H, W = target.shape fp = DEFAULT_FILM_PARAMS sigma_t = torch.tensor([[fp['sigma']]], device=self.device) grain_size_t = torch.tensor([[fp['grain_size']]], device=self.device) lum_params_t = torch.tensor([[fp['lum_a'], fp['lum_b'], fp['lum_c']]], device=self.device) # Halation control: scales the red glow film paints around highlights. # Above 1x the glow color converges on deep film red (the dye layer # that re-exposes when light scatters off the film base), the radius # widens, and the threshold drops so bright skies participate — # giving highlight edges the classic red fringe of color negative. hm = max(0.0, float(halation_mult)) thr = fp['h_threshold'] bias = list(fp['h_color_bias']) if hm > 1.0: t = min((hm - 1.0) / 2.0, 1.0) red_logits = [2.4, -0.7, -2.0] # sigmoid -> [0.92, 0.33, 0.12] bias = [b * (1.0 - t) + r * t for b, r in zip(bias, red_logits)] thr = fp['h_threshold'] - 0.10 * t h_threshold_t = torch.tensor([[thr]], device=self.device) h_radius_t = torch.tensor([[fp['h_radius'] * (1.0 + 0.5 * max(hm - 1.0, 0.0))]], device=self.device) h_intensity_t = torch.tensor([[fp['h_intensity'] * hm]], device=self.device) h_color_bias_t = torch.tensor([bias], device=self.device) # 1. Color grading via predicted LUT graded, predicted_lut, _ = self.style_lut(target, ref_224) # 2. Tone curve matching tone_matched, tone_transfer = match_tone_curve( graded, reference, strength=tone_strength, return_transfer=True) # 3. Film tone (shadow lift, highlight rolloff, color tinting) ref_tone = analyze_reference_tone(reference) film_toned = apply_film_tone(tone_matched, ref_tone, strength=film_tone_strength) # 4. Multi-scale grain grained, _ = render_grain(film_toned, sigma_t, grain_size_t, lum_params_t, grain_mult=grain_mult) # 5. Per-channel halation final, _ = render_halation(grained, h_threshold_t, h_radius_t, h_intensity_t, h_color_bias_t) cube_text = self._lut_to_cube(predicted_lut) params = self._film_params_dict(fp, ref_tone) params.update({ 'tone_transfer_rgb': [ [round(float(v), 6) for v in channel] for channel in tone_transfer.detach().cpu().tolist() ], 'tone_strength': round(float(tone_strength), 4), 'film_tone_strength': round(float(film_tone_strength), 4), 'grain_mult': round(float(grain_mult), 4), 'halation_mult': round(hm, 4), }) curves = self._neutral_curves( predicted_lut, tone_transfer=tone_transfer.detach().cpu(), tone_strength=tone_strength, ref_tone=ref_tone, film_tone_strength=film_tone_strength) ga, gs = self._grain_settings(fp, grain_mult=grain_mult) return { 'final_jpg': self._tensor_to_jpeg_bytes(final), 'graded_jpg': self._tensor_to_jpeg_bytes(graded, quality=90), 'cube_lut': cube_text, 'params': params, 'xmp': self._curves_to_xmp(curves, ga, gs), 'costyle': self._curves_to_costyle(curves, ga, gs), 'size': [W, H], }