3ZadeSSG commited on
Commit
8ff6dec
·
1 Parent(s): bbceeec

initial applicaiton files

Browse files
.gitattributes CHANGED
@@ -33,3 +33,5 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
 
 
 
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
36
+ *.png filter=lfs diff=lfs merge=lfs -text
37
+ *.jpeg filter=lfs diff=lfs merge=lfs -text
.gitignore ADDED
@@ -0,0 +1,16 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Python
2
+ __pycache__/
3
+ *.py[cod]
4
+ *$py.class
5
+
6
+ # Models and Engines
7
+ *.onnx
8
+ *.onnx.data
9
+ *.pth
10
+ *.engine
11
+
12
+ # Videos
13
+ *.mp4
14
+
15
+ # Logs
16
+ logs/
.huggingface.yaml ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ sdk: gradio
2
+ python_version: '3.12'
3
+ requirements_file: requirements.txt
README.md CHANGED
@@ -1,14 +1,14 @@
1
  ---
2
- title: PLFNet
3
- emoji: 💻
4
- colorFrom: pink
5
  colorTo: pink
6
  sdk: gradio
7
  sdk_version: 6.3.0
8
  app_file: app.py
9
  pinned: false
10
  license: agpl-3.0
11
- short_description: Single Image to Real-Time Light Field Reconstruction Model
12
  ---
13
 
14
  Check out the configuration reference at https://huggingface.co/docs/hub/spaces-config-reference
 
1
  ---
2
+ title: PVSNet
3
+ emoji: 🐢
4
+ colorFrom: gray
5
  colorTo: pink
6
  sdk: gradio
7
  sdk_version: 6.3.0
8
  app_file: app.py
9
  pinned: false
10
  license: agpl-3.0
11
+ short_description: Real Time View Synthesis and Light Field Reconstruction Model
12
  ---
13
 
14
  Check out the configuration reference at https://huggingface.co/docs/hub/spaces-config-reference
app.py ADDED
@@ -0,0 +1,474 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import gradio as gr
2
+ import torch
3
+ import numpy as np
4
+ import cv2
5
+ import io
6
+ import tempfile
7
+ import base64
8
+ from PIL import Image
9
+ import torchvision.transforms as transforms
10
+ import parameters as params
11
+ from model import PLFNet
12
+ import helperFunctions as helper
13
+ import socket
14
+ import os
15
+ import json
16
+ from huggingface_hub import hf_hub_download
17
+ import joblib
18
+
19
+ REPO_ID = "3ZadeSSG/PVSNet"
20
+ print("Downloading/Loading checkpoints from Hugging Face Hub...")
21
+ MODEL_FLOWERS_LOCATION = hf_hub_download(
22
+ repo_id=REPO_ID,
23
+ filename="checkpoint_best_flowers.pth"
24
+ )
25
+ MODEL_STANFORD_LOCATION = hf_hub_download(
26
+ repo_id=REPO_ID,
27
+ filename="checkpoint_best_stanford.pth"
28
+ )
29
+
30
+ DEVICE = "cpu"
31
+
32
+ DATASET_CHECKPOINT_MAP = {
33
+ "Flowers": MODEL_FLOWERS_LOCATION,
34
+ "Stanford": MODEL_STANFORD_LOCATION,
35
+ }
36
+
37
+ SAMPLE_IMAGE_DIR = "./sample_images"
38
+ SAMPLE_IMAGES = {}
39
+ for dataset_name in ["Flowers", "Stanford"]:
40
+ folder = os.path.join(SAMPLE_IMAGE_DIR, dataset_name)
41
+ if os.path.isdir(folder):
42
+ images = sorted([
43
+ os.path.join(folder, f)
44
+ for f in os.listdir(folder)
45
+ if f.lower().endswith((".png", ".jpg", ".jpeg"))
46
+ ])
47
+ SAMPLE_IMAGES[dataset_name] = images
48
+
49
+ def getPositionVector(x, y, height, width):
50
+ vector = torch.zeros((2, height, width), dtype=torch.float)
51
+ normalized_x = (x - (-0.003)) / (0.003 - (-0.003))
52
+ normalized_y = (y - (-0.003)) / (0.003 - (-0.003))
53
+ vector[0, :, :] = normalized_x
54
+ vector[1, :, :] = normalized_y
55
+ return vector
56
+
57
+ def predictSingleImage(model, img, target_pose, height, width):
58
+ transform = transforms.Compose([
59
+ transforms.Resize((height, width)),
60
+ transforms.ToTensor()
61
+ ])
62
+ img_input = transform(img).to(DEVICE)
63
+ output_position = getPositionVector(target_pose[0], target_pose[1], height, width).to(DEVICE)
64
+ with torch.no_grad():
65
+ img_ = torch.cat((img_input, output_position), dim=0).unsqueeze(0).to(DEVICE)
66
+ img_out = model(img_).detach().cpu()
67
+ return img_out
68
+
69
+ def generateCircularTrajectory(radius, num_frames, num_loops):
70
+ angles = np.linspace(0, 2 * np.pi * num_loops, num_frames * num_loops)
71
+ return [[radius * np.cos(angle), radius * np.sin(angle)] for angle in angles]
72
+
73
+ def create_video_from_memory(frames, fps=60):
74
+ height, width = frames[0].shape[:2]
75
+ fourcc = cv2.VideoWriter_fourcc(*'mp4v')
76
+ temp_video = tempfile.NamedTemporaryFile(delete=False, suffix=".mp4")
77
+ out = cv2.VideoWriter(temp_video.name, fourcc, fps, (width, height))
78
+ for frame in frames:
79
+ out.write(frame)
80
+ out.release()
81
+ return temp_video.name
82
+
83
+ def process_parallax_video(img, dataset, resolution, radius, num_frames, num_loops):
84
+ if img is None:
85
+ return None
86
+ checkpoint_path = DATASET_CHECKPOINT_MAP.get(dataset, DATASET_CHECKPOINT_MAP["Flowers"])
87
+ model = PLFNet()
88
+ model = helper.load_Checkpoint(checkpoint_path, model, load_cpu=True)
89
+ model.to(DEVICE)
90
+ model.eval()
91
+
92
+ height, width = (352, 512) if "352x512" in resolution else (176, 256)
93
+ img = img.crop((0, 0, img.width, int(img.width * (11 / 16))))
94
+ trajectory = generateCircularTrajectory(radius, num_frames, num_loops)
95
+
96
+ frames = []
97
+ for pose in trajectory:
98
+ output_img = predictSingleImage(model, img, pose, height, width)
99
+ img_np = output_img.squeeze(0).permute(1, 2, 0).numpy()
100
+ img_np = (img_np * 255).astype(np.uint8)
101
+ img_bgr = cv2.cvtColor(img_np, cv2.COLOR_RGB2BGR)
102
+ frames.append(img_bgr)
103
+
104
+ return create_video_from_memory(frames)
105
+
106
+ def generate_lf_raw_frames(img, dataset, resolution):
107
+ if img is None:
108
+ return None, "Please upload an image first."
109
+
110
+ checkpoint_path = DATASET_CHECKPOINT_MAP.get(dataset, DATASET_CHECKPOINT_MAP["Flowers"])
111
+ model = PLFNet()
112
+ model = helper.load_Checkpoint(checkpoint_path, model, load_cpu=True)
113
+ model.to(DEVICE)
114
+ model.eval()
115
+
116
+ height, width = (352, 512) if "352x512" in resolution else (176, 256)
117
+ img = img.crop((0, 0, img.width, int(img.width * (11 / 16))))
118
+
119
+ frames_b64 = []
120
+
121
+ for i in range(-3, 4):
122
+ for j in range(-3, 4):
123
+ pose = [i * 0.001, j * 0.001]
124
+ out = predictSingleImage(model, img, pose, height, width)
125
+ img_np = out.squeeze(0).permute(1, 2, 0).numpy()
126
+ img_np = (img_np * 255).astype(np.uint8)
127
+ img_bgr = cv2.cvtColor(img_np, cv2.COLOR_RGB2BGR)
128
+ _, buffer = cv2.imencode('.jpg', img_bgr, [cv2.IMWRITE_JPEG_QUALITY, 80])
129
+ b64_str = base64.b64encode(buffer).decode('utf-8')
130
+ frames_b64.append(f"data:image/jpeg;base64,{b64_str}")
131
+
132
+ import json
133
+ return json.dumps(frames_b64), "Light Field generated! You can now adjust Focus and Aperture, and move your mouse over the image."
134
+
135
+
136
+
137
+ html_code = """
138
+ <iframe id="lf-iframe" style="width: 100%; max-width: 800px; aspect-ratio: 512/352; border: 2px dashed #ccc; display: block; margin: 0 auto; background: #222;" srcdoc="
139
+ <html>
140
+ <head>
141
+ <script src='https://cdnjs.cloudflare.com/ajax/libs/three.js/r128/three.min.js'></script>
142
+ <script src='https://cdn.jsdelivr.net/npm/three@0.128.0/examples/js/controls/OrbitControls.js'></script>
143
+ <style> body { margin: 0; overflow: hidden; background: #222; } canvas { display: block; width: 100%; height: 100%; } #placeholder { color: #888; font-family: sans-serif; position: absolute; top: 50%; left: 50%; transform: translate(-50%, -50%); pointer-events: none; } #debug { position: absolute; top: 0; left: 0; color: lime; font-family: monospace; padding: 10px; pointer-events: none; z-index: 9999; } </style>
144
+ </head>
145
+ <body>
146
+ <div id='placeholder'>Generated Light Field will appear here</div>
147
+ <div id='debug'></div>
148
+ <script>
149
+ function logDebug(msg) {
150
+ document.getElementById('debug').innerHTML += msg + '<br>';
151
+ }
152
+
153
+ const vertexShader = `
154
+ out vec2 vSt;
155
+ out vec2 vUv;
156
+ void main() {
157
+ vec3 posToCam = cameraPosition - position;
158
+ vec3 nDir = normalize(posToCam);
159
+ float zRatio = posToCam.z / nDir.z;
160
+ vec3 uvPoint = zRatio * nDir;
161
+ vUv = uvPoint.xy + 0.5;
162
+ vUv.x = 1.0 - vUv.x;
163
+ vSt = uv;
164
+ vSt.x = 1.0 - vSt.x;
165
+ gl_Position = projectionMatrix * modelViewMatrix * vec4(position, 1.0);
166
+ }
167
+ `;
168
+
169
+ const fragmentShader = `
170
+ precision highp sampler2DArray;
171
+ uniform sampler2DArray field;
172
+ uniform vec2 camArraySize;
173
+ uniform float aperture;
174
+ uniform float focus;
175
+ in vec2 vSt;
176
+ in vec2 vUv;
177
+ out vec4 fragColor;
178
+
179
+ void main() {
180
+ vec4 color = vec4(0.0);
181
+ float colorCount = 0.0;
182
+ if (vUv.x < 0.0 || vUv.x > 1.0 || vUv.y < 0.0 || vUv.y > 1.0) {
183
+ discard;
184
+ }
185
+ for (float i = 0.0; i < 7.0; i++) {
186
+ for (float j = 0.0; j < 7.0; j++) {
187
+ float dx = i - (vSt.x * camArraySize.x - 0.5);
188
+ float dy = j - (vSt.y * camArraySize.y - 0.5);
189
+ float sqDist = dx * dx + dy * dy;
190
+ if (sqDist <= aperture + 0.001) {
191
+ float camOff = i + camArraySize.x * j;
192
+ vec2 focOff = vec2(dx, dy) * focus;
193
+ color += texture(field, vec3(vUv + focOff, camOff));
194
+ colorCount++;
195
+ }
196
+ }
197
+ }
198
+ fragColor = vec4(color.rgb / max(colorCount, 1.0), 1.0);
199
+ }
200
+ `;
201
+
202
+ let scene, camera, renderer, planeMat, fieldTexture, controls;
203
+ let camsX = 7, camsY = 7;
204
+ let reqFrame;
205
+
206
+ function initScene() {
207
+ if(scene) return;
208
+ scene = new THREE.Scene();
209
+ camera = new THREE.PerspectiveCamera(30, window.innerWidth / window.innerHeight, 0.1, 100);
210
+ camera.position.set(0, 0, 2);
211
+ camera.lookAt(0, 0, -2);
212
+
213
+ const canvas = document.createElement('canvas');
214
+ const context = canvas.getContext('webgl2', { antialias: true });
215
+ if (!context) logDebug('WebGL2 not supported!');
216
+ renderer = new THREE.WebGLRenderer({ canvas: canvas, context: context });
217
+ renderer.setSize(window.innerWidth, window.innerHeight);
218
+ document.body.appendChild(renderer.domElement);
219
+
220
+ controls = new THREE.OrbitControls(camera, renderer.domElement);
221
+ controls.enableDamping = true;
222
+ controls.target.set(0, 0, -2);
223
+ controls.panSpeed = 2;
224
+
225
+ window.addEventListener('resize', () => {
226
+ camera.aspect = window.innerWidth / window.innerHeight;
227
+ camera.updateProjectionMatrix();
228
+ renderer.setSize(window.innerWidth, window.innerHeight);
229
+ });
230
+ }
231
+
232
+ function render() {
233
+ if (controls) controls.update();
234
+ if (renderer && scene && camera) {
235
+ renderer.render(scene, camera);
236
+ }
237
+ reqFrame = requestAnimationFrame(render);
238
+ }
239
+
240
+ window.addEventListener('message', async function(e) {
241
+ if (e.data.type === 'update_frames') {
242
+ document.getElementById('placeholder').innerText = 'Loading textures...';
243
+ document.getElementById('placeholder').style.display = 'block';
244
+
245
+ const frames = e.data.frames;
246
+ if (!frames || frames.length !== 49) {
247
+ logDebug('Error: Invalid frames');
248
+ return;
249
+ }
250
+
251
+ try {
252
+ initScene();
253
+ if (controls) {
254
+ controls.rotateSpeed = e.data.sensitivity;
255
+ controls.panSpeed = e.data.sensitivity * 2;
256
+ }
257
+
258
+ let resX, resY;
259
+ const allBuffer = [];
260
+
261
+ for(let i=0; i<49; i++) {
262
+ const img = new Image();
263
+ await new Promise((resolve, reject) => {
264
+ img.onload = resolve;
265
+ img.onerror = reject;
266
+ img.src = frames[i];
267
+ });
268
+ if(i===0) { resX = img.width; resY = img.height; }
269
+ const cvs = document.createElement('canvas');
270
+ cvs.width = resX; cvs.height = resY;
271
+ const ctx = cvs.getContext('2d', { willReadFrequently: true });
272
+ ctx.drawImage(img, 0, 0);
273
+ const data = ctx.getImageData(0, 0, resX, resY).data;
274
+ allBuffer.push(data);
275
+ }
276
+
277
+ const totalBuffer = new Uint8Array(resX * resY * 4 * 49);
278
+ for(let i=0; i<49; i++) {
279
+ totalBuffer.set(allBuffer[i], i * resX * resY * 4);
280
+ }
281
+
282
+ if(planeMat) {
283
+ scene.remove(scene.children[0]);
284
+ planeMat.dispose();
285
+ }
286
+
287
+ fieldTexture = new THREE.DataTexture2DArray(totalBuffer, resX, resY, 49);
288
+ fieldTexture.format = THREE.RGBAFormat;
289
+ fieldTexture.type = THREE.UnsignedByteType;
290
+ fieldTexture.needsUpdate = true;
291
+
292
+ planeMat = new THREE.ShaderMaterial({
293
+ uniforms: {
294
+ field: { value: fieldTexture },
295
+ camArraySize: { value: new THREE.Vector2(camsX, camsY) },
296
+ aperture: { value: e.data.aperture },
297
+ focus: { value: e.data.focus }
298
+ },
299
+ vertexShader: vertexShader,
300
+ fragmentShader: fragmentShader,
301
+ side: THREE.DoubleSide,
302
+ glslVersion: THREE.GLSL3
303
+ });
304
+
305
+ const planeGeo = new THREE.PlaneGeometry(camsX * 0.1, camsY * 0.1, camsX, camsY);
306
+ const plane = new THREE.Mesh(planeGeo, planeMat);
307
+
308
+ // Scale to correct aspect ratio and increase display size (doesn't break shader parallax logic)
309
+ const aspect = resX / resY;
310
+ plane.scale.set(aspect * 2.0, 2.0, 1.0);
311
+
312
+ plane.position.z = -2;
313
+ scene.add(plane);
314
+
315
+ document.getElementById('placeholder').style.display = 'none';
316
+ if(!reqFrame) render(); // Start animation loop
317
+ } catch (err) {
318
+ logDebug('Error: ' + err.message);
319
+ }
320
+ }
321
+ else if (e.data.type === 'update_sensitivity') {
322
+ if (controls) {
323
+ controls.rotateSpeed = e.data.value;
324
+ controls.panSpeed = e.data.value * 2; // pan is naturally slower
325
+ }
326
+ }
327
+ else if (e.data.type === 'reset_camera') {
328
+ camera.position.set(0, 0, 2);
329
+ camera.up.set(0, 1, 0);
330
+ if (controls) {
331
+ controls.target.set(0, 0, -2);
332
+ controls.update();
333
+ }
334
+ }
335
+ else if (e.data.type === 'update_aperture') {
336
+ if(planeMat) planeMat.uniforms.aperture.value = e.data.value;
337
+ }
338
+ });
339
+ </script>
340
+ </body>
341
+ </html>
342
+ "></iframe>
343
+ """
344
+
345
+ def load_sample_image(image_path, dataset_name):
346
+ img = Image.open(image_path)
347
+ return img, dataset_name
348
+
349
+ with gr.Blocks(title="PVSNet/PLFNet", theme="default") as demo:
350
+ gr.Markdown("""
351
+ # PVSNet: Real-Time Position-Aware View Synthesis from Single-View Input
352
+ * Upload a single Lytro image and get a mini parallax video showing capabilities of the light field reconstruction model from our works PVSNet and PLFNet.
353
+ **Note** Huggingface demo is running on CPU, so the inference speed will be slow. It might take around 2 minutes for full LF reconstruction or video generation.
354
+ ### Head to our [Project Page](https://realistic3d-miun.github.io/PVSNet/) for more details about the models.
355
+ """)
356
+
357
+ with gr.Row():
358
+ img_input = gr.Image(type="pil", label="Upload Image")
359
+ with gr.Column():
360
+ dataset = gr.Dropdown(
361
+ choices=["Flowers", "Stanford"],
362
+ value="Flowers",
363
+ label="Model Checkpoint (Dataset)"
364
+ )
365
+ resolution = gr.Dropdown(["352x512 (Slow)", "176x256 (Fast)"], value="352x512 (Slow)", label="Resolution")
366
+
367
+ with gr.Tabs():
368
+ with gr.Tab("Parallax Video"):
369
+ with gr.Row():
370
+ with gr.Column():
371
+ radius = gr.Slider(0.0006, 0.006, value=0.003, label="Radius")
372
+ num_frames = gr.Slider(10, 100, value=60, step=10, label="Number of Frames")
373
+ num_loops = gr.Slider(1, 10, value=4, step=1, label="Number of Loops")
374
+ generate_vid_btn = gr.Button("Generate Video", variant="primary")
375
+ video_output = gr.Video(label="Generated Video", height=352)
376
+
377
+ generate_vid_btn.click(
378
+ fn=process_parallax_video,
379
+ inputs=[img_input, dataset, resolution, radius, num_frames, num_loops],
380
+ outputs=video_output,
381
+ )
382
+
383
+ with gr.Tab("Interactive Light Field"):
384
+ with gr.Row():
385
+ with gr.Column():
386
+ generate_lf_btn = gr.Button("Generate Light Field Data", variant="primary")
387
+ lf_status = gr.Textbox(label="Status", interactive=False, value="Awaiting generation...")
388
+
389
+ gr.Markdown("### Rendering Controls\nAdjust these parameters and use your mouse to navigate the Light Field (Left Click: Rotate, Right Click: Pan, Scroll: Zoom).")
390
+ with gr.Row():
391
+ sensitivity = gr.Slider(0.1, 3.0, value=1.0, step=0.1, label="Mouse Sensitivity")
392
+ aperture = gr.Slider(0, 4.5, value=2.2, step=0.1, label="Aperture")
393
+ reset_btn = gr.Button("Reset Camera")
394
+
395
+ with gr.Column():
396
+ gr.HTML(html_code)
397
+
398
+ # Hidden text box to transfer JSON data to frontend
399
+ b64_frames_state = gr.Textbox(visible=False, elem_id="lf_data_bridge")
400
+
401
+ # JS function string to update frames in iframe
402
+ update_js = """(val, a, s) => {
403
+ if (val) {
404
+ const frames = JSON.parse(val);
405
+ const iframe = document.getElementById('lf-iframe');
406
+ if (iframe && iframe.contentWindow) {
407
+ iframe.contentWindow.postMessage({
408
+ type: 'update_frames',
409
+ frames: frames,
410
+ aperture: a,
411
+ focus: 0,
412
+ sensitivity: s
413
+ }, '*');
414
+ }
415
+ }
416
+ }"""
417
+
418
+ # Step 1: Generate Raw Light Field (7x7 array)
419
+ generate_lf_btn.click(
420
+ fn=generate_lf_raw_frames,
421
+ inputs=[img_input, dataset, resolution],
422
+ outputs=[b64_frames_state, lf_status]
423
+ ).then( # Step 2: Run JS to update frontend
424
+ fn=None,
425
+ inputs=[b64_frames_state, aperture, sensitivity],
426
+ outputs=None,
427
+ js=update_js
428
+ )
429
+
430
+ # Sliders update iframe instantly via postMessage (no Python execution needed!)
431
+ sensitivity.change(
432
+ fn=None,
433
+ inputs=[sensitivity],
434
+ outputs=None,
435
+ js="(s) => { const iframe = document.getElementById('lf-iframe'); if (iframe && iframe.contentWindow) iframe.contentWindow.postMessage({type: 'update_sensitivity', value: s}, '*'); }"
436
+ )
437
+
438
+ aperture.change(
439
+ fn=None,
440
+ inputs=[aperture],
441
+ outputs=None,
442
+ js="(a) => { const iframe = document.getElementById('lf-iframe'); if (iframe && iframe.contentWindow) iframe.contentWindow.postMessage({type: 'update_aperture', value: a}, '*'); }"
443
+ )
444
+
445
+ reset_btn.click(
446
+ fn=None,
447
+ inputs=None,
448
+ outputs=None,
449
+ js="() => { const iframe = document.getElementById('lf-iframe'); if (iframe && iframe.contentWindow) iframe.contentWindow.postMessage({type: 'reset_camera'}, '*'); }"
450
+ )
451
+
452
+ gr.Markdown("### Example Images: Click to Load")
453
+ for dataset_name, images in SAMPLE_IMAGES.items():
454
+ with gr.Accordion(f"📂 {dataset_name} Samples", open=(dataset_name == "Flowers")):
455
+ for row_start in range(0, len(images), 3):
456
+ row_images = images[row_start : row_start + 3]
457
+ with gr.Row():
458
+ for img_path in row_images:
459
+ label = os.path.splitext(os.path.basename(img_path))[0]
460
+ sample_img = gr.Image(
461
+ img_path,
462
+ label=label,
463
+ height=150,
464
+ interactive=False,
465
+ show_label=True,
466
+ )
467
+ sample_img.select(
468
+ fn=lambda path=img_path, ds=dataset_name: load_sample_image(path, ds),
469
+ inputs=[],
470
+ outputs=[img_input, dataset],
471
+ )
472
+
473
+ if __name__ == "__main__":
474
+ demo.launch()
helperFunctions.py ADDED
@@ -0,0 +1,26 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import os
3
+ import torch.nn.functional as F
4
+
5
+ def save_checkpoint(model, filelocation, save_parallel = True):
6
+ if save_parallel:
7
+ torch.save(model.module.state_dict(), filelocation)
8
+ else:
9
+ torch.save(model.state_dict(), filelocation)
10
+
11
+ def load_Checkpoint(fileLocation,model, load_cpu=False):
12
+ if load_cpu:
13
+ model.load_state_dict(torch.load(fileLocation,map_location=lambda storage, loc: storage))
14
+ else:
15
+ model.load_state_dict(torch.load(fileLocation))
16
+ return model
17
+
18
+ def writeLog(logList, filename):
19
+ with open(filename, 'w') as outfile:
20
+ outfile.write("\n".join(logList))
21
+
22
+
23
+ def kl_loss(mu, logvar):
24
+ return -0.5 * (1 + logvar - mu.pow(2) - logvar.exp()).mean()
25
+
26
+
model.py ADDED
@@ -0,0 +1,263 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import torch.nn as nn
3
+ import torch.nn.functional as F
4
+ import warnings
5
+ warnings.filterwarnings("ignore")
6
+ import torchvision
7
+ import parameters as params
8
+
9
+ def getLinearLayer(in_feat, out_feat, activation=nn.ReLU(True)):
10
+ return nn.Sequential(
11
+ nn.Linear(in_features=in_feat, out_features=out_feat, bias=True),
12
+ activation
13
+ )
14
+
15
+ def getConvLayer(in_channel,out_channel,stride=1,padding=1,activation=nn.ReLU()):
16
+ return nn.Sequential(nn.Conv2d(in_channel,
17
+ out_channel,
18
+ kernel_size=3,
19
+ stride=stride,
20
+ padding=padding,
21
+ padding_mode='reflect'),
22
+ activation)
23
+
24
+ def getConvTransposeLayer(in_channel, out_channel,kernel=3,stride=1,padding=1,activation=nn.ReLU()):
25
+ return nn.Sequential(nn.ConvTranspose2d(in_channel,
26
+ out_channel,
27
+ kernel_size = kernel,
28
+ stride=stride,
29
+ padding=padding),
30
+ activation)
31
+
32
+
33
+
34
+ class Flatten(nn.Module):
35
+ def forward(self, input):
36
+ return input.view(input.size(0), -1)
37
+
38
+ class UnFlatten(nn.Module):
39
+ def forward(self, input, size=1):
40
+ return input.view(input.size(0), 1, params.params_height//16, params.params_width//16)
41
+
42
+ class ResidualBlock(nn.Module):
43
+ def __init__(self, in_channels, out_channels, stride=1):
44
+ super(ResidualBlock, self).__init__()
45
+ self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=stride, padding=1, bias=False)
46
+ self.relu = nn.ReLU()
47
+ self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1, bias=False)
48
+ self.stride = stride
49
+
50
+ self.shortcut = nn.Sequential()
51
+ if stride != 1 or in_channels != out_channels:
52
+ self.shortcut = nn.Sequential(
53
+ nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=stride, bias=False),
54
+ nn.BatchNorm2d(out_channels)
55
+ )
56
+
57
+ def forward(self, x):
58
+ residual = x
59
+
60
+ out = self.conv1(x)
61
+ out = self.relu(out)
62
+
63
+ out = self.conv2(out)
64
+
65
+ out = out + self.shortcut(residual)
66
+ out = self.relu(out)
67
+ return out
68
+
69
+ class MLPEncoder(nn.Module):
70
+ def __init__(self):
71
+ super().__init__()
72
+
73
+ self.flat = Flatten()
74
+ self.layer1 = getLinearLayer((params.params_height//8)*(params.params_width//8)*2, 1024)
75
+ self.layer2 = getLinearLayer(1024, 512)
76
+ self.layer3 = getLinearLayer(512, (params.params_height//16)*(params.params_width//16))
77
+ self.unflat = UnFlatten()
78
+ self.up_layer1 = nn.Upsample(scale_factor=2, mode='nearest')
79
+ self.up_layer2 = nn.Upsample(scale_factor=2, mode='nearest')
80
+ self.up_layer3 = nn.Upsample(scale_factor=2, mode='nearest')
81
+ self.up_layer4 = nn.Upsample(scale_factor=2, mode='nearest')
82
+
83
+ def forward(self, x):
84
+ x = self.flat(x)
85
+
86
+ x = self.layer1(x)
87
+ x = self.layer2(x)
88
+ x = self.layer3(x)
89
+
90
+ x = self.unflat(x)
91
+
92
+ x = self.up_layer1(x)
93
+ x = self.up_layer2(x)
94
+ x = self.up_layer3(x)
95
+ x = self.up_layer4(x)
96
+ return x
97
+
98
+ class UpperEncoder(nn.Module):
99
+ def __init__(self):
100
+ super().__init__()
101
+ model = torchvision.models.resnet152(pretrained=False)
102
+ layers = list(model.children())
103
+ self.ResNetEncoder = torch.nn.Sequential(*layers[:5].copy())
104
+ del model
105
+
106
+ def forward(self, x):
107
+ x1 = x[:, 0:3, :, :]
108
+ x1 = self.ResNetEncoder(x1)
109
+ return x1
110
+
111
+ def apply_resnet_encoder(self, x):
112
+ x1 = x[:, 0:3, :, :]
113
+ x1 = self.ResNetEncoder(x1)
114
+ return x1
115
+
116
+ class LowerEncoder(nn.Module):
117
+ def __init__(self,total_image_input=1):
118
+ super().__init__()
119
+ self.encoder_pre = ResidualBlock((total_image_input*3)+2, 20)
120
+ self.encoder_layer1 = ResidualBlock(20, 30)
121
+ self.encoder_layer2 = ResidualBlock(30, 50)
122
+
123
+ self.encoder_layer3 = nn.Sequential(
124
+ ResidualBlock(50, 100),
125
+ nn.MaxPool2d(kernel_size=2, stride=2)
126
+ )
127
+
128
+ self.encoder_layer4 = ResidualBlock(100, 200)
129
+ self.encoder_layer5 = nn.Sequential(
130
+ ResidualBlock(200, 400),
131
+ nn.MaxPool2d(kernel_size=2, stride=2)
132
+ )
133
+
134
+ self.encoder_layer6 = ResidualBlock(400, 600)
135
+ self.encoder_layer7 = nn.Sequential(
136
+ ResidualBlock(600, 800),
137
+ nn.MaxPool2d(kernel_size=2, stride=2)
138
+ )
139
+
140
+ self.encoder_layer8 = ResidualBlock(800, 1000)
141
+ self.encoder_layer9 = nn.Sequential(
142
+ ResidualBlock(1000, 1200),
143
+ nn.MaxPool2d(kernel_size=2, stride=2)
144
+ )
145
+
146
+ self.encoder_layer10 = ResidualBlock(1200, 1400)
147
+ self.encoder_layer11 = ResidualBlock(1400, 1600)
148
+
149
+ def forward(self, x):
150
+ x = self.encoder_pre(x)
151
+ x = self.encoder_layer1(x)
152
+ x = self.encoder_layer2(x)
153
+ skip1 = self.encoder_layer3(x)
154
+
155
+ x = self.encoder_layer4(skip1)
156
+ skip2 = self.encoder_layer5(x)
157
+
158
+ x = self.encoder_layer6(skip2)
159
+ skip3 = self.encoder_layer7(x)
160
+
161
+ x = self.encoder_layer8(skip3)
162
+ skip4 = self.encoder_layer9(x)
163
+
164
+ x = self.encoder_layer10(skip4)
165
+ x = self.encoder_layer11(x)
166
+
167
+ return x, [skip1, skip2, skip3, skip4]
168
+
169
+ class MergeDecoder(nn.Module):
170
+ def __init__(self):
171
+ super().__init__()
172
+
173
+ self.decoder_layer1 = ResidualBlock(1600, 1400)
174
+ self.decoder_layer2 = ResidualBlock(1400, 1200)
175
+ self.decoder_layer3 = ResidualBlock(1200, 1000)
176
+
177
+ self.decoder_layer4 = nn.Sequential(
178
+ nn.ConvTranspose2d(1000, 800, 2, stride=2, padding=0),
179
+ nn.ReLU(True)
180
+ )
181
+ self.decoder_layer5 = ResidualBlock(800, 600)
182
+
183
+ self.decoder_layer6 = nn.Sequential(
184
+ nn.ConvTranspose2d(600, 400, 2, stride=2, padding=0),
185
+ nn.ReLU(True)
186
+ )
187
+ self.decoder_layer7 = ResidualBlock(400, 200)
188
+
189
+ self.decoder_layer8 = nn.Sequential(
190
+ nn.ConvTranspose2d(200, 100, 2, stride=2, padding=0),
191
+ nn.ReLU(True)
192
+ )
193
+ self.decoder_layer9 = ResidualBlock(100, 100)
194
+
195
+ self.decoder_layer10 = nn.Sequential(
196
+ nn.ConvTranspose2d(100, 100, 2, stride=2, padding=0),
197
+ nn.ReLU(True)
198
+ )
199
+ self.decoder_layer11 = ResidualBlock(100, 100)
200
+ self.decoder_layer12 = ResidualBlock(100, 50)
201
+ self.decoder_layer13 = ResidualBlock(50, 40)
202
+ self.decoder_layer14 = ResidualBlock(40, 20)
203
+ self.decoder_layer15 = nn.Sequential(
204
+ nn.Conv2d(20, 8, 3, stride=1, padding=1),
205
+ nn.Sigmoid()
206
+ )
207
+ self.decoder_layer16 = nn.Sequential(
208
+ nn.Conv2d(8, 3, 3, stride=1, padding=1),
209
+ nn.Sigmoid()
210
+ )
211
+
212
+ def forward(self, x, lower_skip_list, upper_skip_list):
213
+ x = self.decoder_layer1(x)
214
+ x = self.decoder_layer2(x)
215
+ x = x + lower_skip_list[3] + upper_skip_list[1]
216
+
217
+ x = self.decoder_layer3(x)
218
+ x = self.decoder_layer4(x)
219
+ x = x + lower_skip_list[2] + upper_skip_list[0]
220
+
221
+ x = self.decoder_layer5(x)
222
+ x = self.decoder_layer6(x)
223
+ x = x + lower_skip_list[1]
224
+
225
+ x = self.decoder_layer7(x)
226
+ x = self.decoder_layer8(x)
227
+ x = x + lower_skip_list[0]
228
+
229
+ x = self.decoder_layer9(x)
230
+ x = self.decoder_layer10(x)
231
+ x = self.decoder_layer11(x)
232
+ x = self.decoder_layer12(x)
233
+ x = self.decoder_layer13(x)
234
+ x = self.decoder_layer14(x)
235
+ x = self.decoder_layer15(x)
236
+ x = self.decoder_layer16(x)
237
+ return x
238
+
239
+ class PLFNet(nn.Module):
240
+ def __init__(self,total_image_input=1):
241
+ super().__init__()
242
+ self.upper_encoder = UpperEncoder()
243
+ self.lower_encoder = LowerEncoder(total_image_input)
244
+ self.merge_decoder = MergeDecoder()
245
+
246
+ self.upper_encoder_extra_1 = nn.Sequential(
247
+ ResidualBlock(256, 800),
248
+ nn.MaxPool2d(kernel_size=2, stride=2)
249
+ )
250
+ self.upper_encoder_extra_2 = nn.Sequential(
251
+ ResidualBlock(800, 1200),
252
+ nn.MaxPool2d(kernel_size=2, stride=2)
253
+ )
254
+
255
+ def forward(self, x):
256
+ upper_features_1 = self.upper_encoder.apply_resnet_encoder(x)
257
+ upper_features_1 = self.upper_encoder_extra_1(upper_features_1)
258
+ upper_features_2 = self.upper_encoder_extra_2(upper_features_1)
259
+
260
+ lower_feature, skip_list = self.lower_encoder(x)
261
+ merged_feature = self.merge_decoder(lower_feature, skip_list, [upper_features_1, upper_features_2])
262
+
263
+ return merged_feature
parameters.py ADDED
@@ -0,0 +1,25 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+
3
+ params_width = 512
4
+ params_height = 352
5
+
6
+ TRAIN_LOCATION = "./lf_train.txt"
7
+ VALIDATION_LOCATION = "./lf_validate.txt"
8
+ TEST_LOCATION = "./lf_test.txt"
9
+ LOG_FILE_LOCATION = "./logs/training_log_0.txt"
10
+ CHECKPOINT_LOCATION = "./checkpoint/"
11
+ RESUME_CHECKPOINT_LOCATION = "./checkpoint/checkpoint_best.pth"
12
+ START_CHECKPOINT_LOCATION = "./checkpoint/checkpoint_init.pth"
13
+ DEVICE = "cpu"
14
+
15
+ BATCH_SIZE = 16
16
+ LEARNING_RATE = 0.0001
17
+ NUM_EPOCHS = 150
18
+ START_EPOCH = 0
19
+ PRINT_INTERVAL = 20
20
+
21
+ os.makedirs("./logs",exist_ok=True)
22
+ os.makedirs("./checkpoint",exist_ok=True)
23
+ os.makedirs("./output",exist_ok=True)
24
+
25
+
requirements.txt ADDED
@@ -0,0 +1,18 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ numpy
2
+ torch==2.9.1
3
+ torchvision==0.24.1
4
+ pytorch-msssim==1.0.0
5
+ pytorchvideo==0.1.5
6
+ gradio==6.2.0
7
+ gradio_client==2.0.2
8
+ opencv-python==4.6.0.66
9
+ pillow==10.4.0
10
+ pillow_heif==0.15.0
11
+ matplotlib==3.10.8
12
+ matplotlib-inline==0.1.6
13
+ tqdm==4.65.0
14
+ moviepy==1.0.3
15
+ scikit-image==0.26.0
16
+ scikit-learn==1.8.0
17
+ scipy==1.11.4
18
+ random-fourier-features-pytorch
sample_images/Flowers/104_image_3_3.png ADDED

Git LFS Details

  • SHA256: 808bb5000fb79900f8e498d2071e579b07558b55f5f227f0a3cd2ec8f4b3e934
  • Pointer size: 131 Bytes
  • Size of remote file: 241 kB
sample_images/Flowers/193_image_3_3.png ADDED

Git LFS Details

  • SHA256: 9e99b7cfc7a523d5f33a5eff2626ead640a0c5fc36cfb3353b4e4a9063846902
  • Pointer size: 131 Bytes
  • Size of remote file: 247 kB
sample_images/Flowers/20_image_3_3.png ADDED

Git LFS Details

  • SHA256: 265ed2d3a4cdd79baf9bac12d228f52cdcefab6437920b0a9e44f43dca413226
  • Pointer size: 131 Bytes
  • Size of remote file: 289 kB
sample_images/Flowers/28_image_3_3.png ADDED

Git LFS Details

  • SHA256: c2f57e6906c0bc3e59bfdf0a2946e49930653b4e6f787a43fb265698139a04c3
  • Pointer size: 131 Bytes
  • Size of remote file: 244 kB
sample_images/Flowers/320_image_3_3.png ADDED

Git LFS Details

  • SHA256: 7cd35487cbd7e5014ee1184cd6b224dfb306387fdafe2536f3f9f15dc372be13
  • Pointer size: 131 Bytes
  • Size of remote file: 234 kB
sample_images/Flowers/321_image_3_3.png ADDED

Git LFS Details

  • SHA256: ee23dfdd3e02a60383e46aca118b15b09b9fb34a678c04082c1bb89222f44e6d
  • Pointer size: 131 Bytes
  • Size of remote file: 211 kB
sample_images/Flowers/44_image_3_3.png ADDED

Git LFS Details

  • SHA256: 391d2eeca47d9ed5cd3fbf4a0b17cfe4bda11902f5aa0c3335ecddcb572ad5ab
  • Pointer size: 131 Bytes
  • Size of remote file: 279 kB
sample_images/Flowers/uploaded_image.png ADDED

Git LFS Details

  • SHA256: 7cd35487cbd7e5014ee1184cd6b224dfb306387fdafe2536f3f9f15dc372be13
  • Pointer size: 131 Bytes
  • Size of remote file: 234 kB
sample_images/Stanford/106_image_3_3.png ADDED

Git LFS Details

  • SHA256: 865f1ea85d77e377c4f166dc4ba1c97ed4f742fa2fc9a94d8d44a4b392fcb8e3
  • Pointer size: 131 Bytes
  • Size of remote file: 269 kB
sample_images/Stanford/166_image_3_3.png ADDED

Git LFS Details

  • SHA256: 5decf6cacf467dfd4157763b8ce0f08f6dd9d0be95fa996026198b6bef5efae1
  • Pointer size: 131 Bytes
  • Size of remote file: 241 kB
sample_images/Stanford/183_image_3_3.png ADDED

Git LFS Details

  • SHA256: 35719aeb58dbadb58e83f0cce614fa250a6885cca4dd30f37574ff9f8cc91833
  • Pointer size: 131 Bytes
  • Size of remote file: 209 kB
sample_images/Stanford/185_image_3_3.png ADDED

Git LFS Details

  • SHA256: 448396c32c7666fd20c682931cd9f6ef1645e511559efd026bb960e00c493db7
  • Pointer size: 131 Bytes
  • Size of remote file: 278 kB
sample_images/Stanford/18_image_3_3.png ADDED

Git LFS Details

  • SHA256: f7f19cef480a286cfe8b2dc4bf13a30629bcdfbfefce02c3b586cc0676afd710
  • Pointer size: 131 Bytes
  • Size of remote file: 305 kB
sample_images/uploaded_image.png ADDED

Git LFS Details

  • SHA256: 9e99b7cfc7a523d5f33a5eff2626ead640a0c5fc36cfb3353b4e4a9063846902
  • Pointer size: 131 Bytes
  • Size of remote file: 247 kB