import os import sys import pty import select import signal import subprocess import threading import time from flask import Flask, jsonify, request, render_template, send_from_directory, Response app = Flask(__name__, template_folder='templates', static_folder='static') # Global process variables current_process = None master_fd = None log_buffer = "" running_script = None log_lock = threading.Lock() # Define script workflow categories SCRIPT_CATEGORIES = { "workspace": { "title": "Workspace & Setup", "icon": "folder", "scripts": [ {"file": "1_clear_workspace.sh", "name": "Clear Workspace", "desc": "Deletes the workspace directory and all intermediate files to start clean."} ] }, "src_video": { "title": "Step 1: Source Video (data_src)", "icon": "video", "scripts": [ {"file": "2_extract_image_from_data_src.sh", "name": "Extract Images", "desc": "Extracts video frames from workspace/data_src.mp4 to PNG files."}, {"file": "4_data_src_extract_faces_S3FD.sh", "name": "Extract Faces (S3FD)", "desc": "Uses S3FD detector to extract aligned face crops from source frames."}, {"file": "4_data_src_extract_faces_MANUAL.sh", "name": "Extract Faces (Manual)", "desc": "Interactively extract or fix missing face detections manually."}, {"file": "4.2_data_src_sort.sh", "name": "Sort Source Faces", "desc": "Sorts source faces by similarity, luminance, blur, etc., to prune trash."} ] }, "dst_video": { "title": "Step 2: Destination Video (data_dst)", "icon": "video", "scripts": [ {"file": "3_extract_image_from_data_dst.sh", "name": "Extract Images", "desc": "Extracts video frames from workspace/data_dst.mp4 to PNG files."}, {"file": "3.1_denoise_data_dst_images.sh", "name": "Denoise Images", "desc": "Denoises the extracted destination frames for cleaner final merge."}, {"file": "5_data_dst_extract_faces_S3FD.sh", "name": "Extract Faces (S3FD)", "desc": "Uses S3FD detector to extract aligned faces from destination frames."}, {"file": "5_data_dst_extract_faces_S3FD_+_manual_fix.sh", "name": "Extract Faces (S3FD + Manual)", "desc": "Extract faces using S3FD with manual correction fallback."}, {"file": "5_data_dst_extract_faces_MANUAL.sh", "name": "Extract Faces (Manual)", "desc": "Interactively extract or fix destination faces manually."} ] }, "xseg": { "title": "Step 3: XSeg Masking (Optional)", "icon": "scissors", "scripts": [ {"file": "5_XSeg_train.sh", "name": "Train XSeg Mask", "desc": "Trains a custom XSeg masking model on labeled faces."}, {"file": "5_XSeg_data_src_mask_edit.sh", "name": "Edit Src Masks", "desc": "Label/edit XSeg masks on your source face images."}, {"file": "5_XSeg_data_src_mask_apply.sh", "name": "Apply Src Masks", "desc": "Applies a trained XSeg model to mask source face images."}, {"file": "5_XSeg_data_dst_mask_edit.sh", "name": "Edit Dst Masks", "desc": "Label/edit XSeg masks on your destination face images."}, {"file": "5_XSeg_data_dst_mask_apply.sh", "name": "Apply Dst Masks", "desc": "Applies a trained XSeg model to mask destination face images."} ] }, "training": { "title": "Step 4: Model Training", "icon": "cpu", "scripts": [ {"file": "6_train_Quick96_no_preview.sh", "name": "Train Quick96 (Headless)", "desc": "Trains Quick96 model in the background (no local window)."}, {"file": "6_train_Quick96.sh", "name": "Train Quick96 (Windowed)", "desc": "Trains Quick96 model. Requires local X11 visual window."}, {"file": "6_train_SAEHD_no_preview.sh", "name": "Train SAEHD (Headless)", "desc": "Trains high-quality SAEHD model in the background."}, {"file": "6_train_SAEHD.sh", "name": "Train SAEHD (Windowed)", "desc": "Trains SAEHD model. Requires local X11 visual window."} ] }, "merging": { "title": "Step 5: Face Merging", "icon": "merge", "scripts": [ {"file": "7_merge_Quick96.sh", "name": "Merge Quick96", "desc": "Merges Quick96 face swaps onto destination frames."}, {"file": "7_merge_SAEHD.sh", "name": "Merge SAEHD", "desc": "Merges SAEHD face swaps onto destination frames."} ] }, "export": { "title": "Step 6: Export & Render", "icon": "download", "scripts": [ {"file": "8_merged_to_mp4.sh", "name": "Export to MP4", "desc": "Combines merged frames into an MP4 file (h264)."}, {"file": "8_merged_to_mp4_lossless.sh", "name": "Export to MP4 (Lossless)", "desc": "Combines merged frames into a lossless MP4 file."}, {"file": "8_merged_to_avi.sh", "name": "Export to AVI", "desc": "Combines merged frames into an uncompressed AVI file."} ] } } # Add pre-trained download links to categories SCRIPT_CATEGORIES["pretrain"] = { "title": "Download Pre-trained Weights", "icon": "cloud-download", "scripts": [ {"file": "4.1_download_Quick96.sh", "name": "Download Quick96 pretrain", "desc": "Downloads Quick96 pretraining zip file."}, {"file": "4.1_download_CelebA.sh", "name": "Download CelebA pretrain", "desc": "Downloads CelebA pretraining zip file."}, {"file": "4.1_download_FFHQ.sh", "name": "Download FFHQ pretrain", "desc": "Downloads FFHQ pretraining zip file."} ] } def log_reader(fd, proc): global log_buffer, current_process, master_fd, running_script while True: try: r, w, x = select.select([fd], [], [], 0.1) if fd in r: data = os.read(fd, 4096) if not data: break decoded_chunk = data.decode('utf-8', errors='replace') with log_lock: log_buffer += decoded_chunk except (OSError, ValueError): break except Exception: break # Check if process died if proc.poll() is not None: # Final drain of buffer try: r, w, x = select.select([fd], [], [], 0.1) if fd in r: data = os.read(fd, 4096) if data: decoded_chunk = data.decode('utf-8', errors='replace') with log_lock: log_buffer += decoded_chunk except Exception: pass break try: os.close(fd) except Exception: pass with log_lock: if current_process == proc: current_process = None master_fd = None running_script = None def get_gpu_info(): try: result = subprocess.run( ['nvidia-smi', '--query-gpu=gpu_name,utilization.gpu,utilization.memory,memory.total,memory.used,temperature.gpu', '--format=csv,noheader,nounits'], stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True, timeout=2 ) if result.returncode == 0: lines = result.stdout.strip().split('\n') gpus = [] for line in lines: parts = [p.strip() for p in line.split(',')] if len(parts) >= 6: gpus.append({ 'name': parts[0], 'gpu_util': parts[1] + '%', 'mem_util': parts[2] + '%', 'mem_total': parts[3] + ' MB', 'mem_used': parts[4] + ' MB', 'temp': parts[5] + '°C' }) return gpus except Exception: pass return None def get_system_stats(): # Basic system utilization (cpu / memory) # Reading from /proc/stat cpu_util = "0%" mem_util = "0%" try: # CPU calculation with open('/proc/stat', 'r') as f: fields = [float(column) for column in f.readline().strip().split()[1:]] idle, total = fields[3], sum(fields) time.sleep(0.1) with open('/proc/stat', 'r') as f: fields2 = [float(column) for column in f.readline().strip().split()[1:]] idle2, total2 = fields2[3], sum(fields2) diff_idle = idle2 - idle diff_total = total2 - total if diff_total > 0: cpu_util = f"{int((1.0 - diff_idle / diff_total) * 100)}%" # Memory calculation with open('/proc/meminfo', 'r') as f: lines = f.readlines() mem_total = 1.0 mem_free = 0.0 mem_cached = 0.0 mem_buffers = 0.0 for line in lines: if line.startswith('MemTotal:'): mem_total = float(line.split()[1]) elif line.startswith('MemFree:'): mem_free = float(line.split()[1]) elif line.startswith('Cached:'): mem_cached = float(line.split()[1]) elif line.startswith('Buffers:'): mem_buffers = float(line.split()[1]) mem_used = mem_total - mem_free - mem_cached - mem_buffers mem_util = f"{int((mem_used / mem_total) * 100)}%" except Exception: pass return { 'cpu': cpu_util, 'mem': mem_util } @app.route('/') def index(): return render_template('index.html') @app.route('/api/status') def status(): global current_process, running_script is_running = current_process is not None and current_process.poll() is None # Get workspace stats workspace_dir = os.path.join(os.path.dirname(__file__), 'workspace') stats = { 'src_frames': 0, 'src_faces': 0, 'dst_frames': 0, 'dst_faces': 0, 'models': [] } if os.path.exists(workspace_dir): def count_files(path, ext=None): if not os.path.exists(path): return 0 count = 0 for entry in os.scandir(path): if entry.is_file(): if ext is None or entry.name.lower().endswith(ext): count += 1 return count stats['src_frames'] = count_files(os.path.join(workspace_dir, 'data_src'), '.png') + count_files(os.path.join(workspace_dir, 'data_src'), '.jpg') stats['src_faces'] = count_files(os.path.join(workspace_dir, 'data_src', 'aligned'), '.jpg') stats['dst_frames'] = count_files(os.path.join(workspace_dir, 'data_dst'), '.png') + count_files(os.path.join(workspace_dir, 'data_dst'), '.jpg') stats['dst_faces'] = count_files(os.path.join(workspace_dir, 'data_dst', 'aligned'), '.jpg') model_dir = os.path.join(workspace_dir, 'model') if os.path.exists(model_dir): for entry in os.scandir(model_dir): if entry.is_file() and entry.name.endswith('.dat'): stats['models'].append(entry.name) return jsonify({ 'running': is_running, 'script': running_script, 'gpu': get_gpu_info(), 'system': get_system_stats(), 'workspace': stats }) @app.route('/api/scripts') def list_scripts(): return jsonify(SCRIPT_CATEGORIES) @app.route('/api/run', methods=['POST']) def run_script(): global current_process, master_fd, log_buffer, running_script if current_process is not None and current_process.poll() is None: return jsonify({'error': 'A script is already running'}), 400 data = request.json script_name = data.get('script') if not script_name: return jsonify({'error': 'No script name specified'}), 400 script_path = os.path.join(os.path.dirname(__file__), 'scripts', script_name) if not os.path.exists(script_path): return jsonify({'error': f'Script not found: {script_name}'}), 404 # Reset log buffer with log_lock: log_buffer = f"=== Starting {script_name} ===\n" running_script = script_name # Spawn the script in a pseudo-terminal try: m_fd, s_fd = pty.openpty() # We spawn the script inside the conda env's python path by launching /bin/bash in pty current_process = subprocess.Popen( ['/bin/bash', script_name], cwd=os.path.join(os.path.dirname(__file__), 'scripts'), stdin=s_fd, stdout=s_fd, stderr=s_fd, preexec_fn=os.setsid, # Put subprocess in its own process group to kill children env=os.environ.copy() ) # Close the slave file descriptor in the parent os.close(s_fd) master_fd = m_fd # Start background reader thread t = threading.Thread(target=log_reader, args=(m_fd, current_process)) t.daemon = True t.start() return jsonify({'status': 'started', 'script': script_name}) except Exception as e: with log_lock: log_buffer += f"\nFailed to launch script: {str(e)}\n" current_process = None master_fd = None running_script = None return jsonify({'error': f'Launch error: {str(e)}'}), 500 @app.route('/api/stop', methods=['POST']) def stop_script(): global current_process if current_process is None or current_process.poll() is not None: return jsonify({'error': 'No script is currently running'}), 400 try: # Kill the entire process group pgid = os.getpgid(current_process.pid) os.killpg(pgid, signal.SIGTERM) time.sleep(0.5) if current_process.poll() is None: os.killpg(pgid, signal.SIGKILL) with log_lock: log_buffer += "\n=== Process terminated by user ===\n" return jsonify({'status': 'terminated'}) except Exception as e: return jsonify({'error': f'Error terminating process: {str(e)}'}), 500 @app.route('/api/input', methods=['POST']) def send_input(): global master_fd data = request.json text = data.get('text', '') if master_fd is None: return jsonify({'error': 'No active process terminal to write to'}), 400 try: os.write(master_fd, (text + '\n').encode('utf-8')) return jsonify({'status': 'sent'}) except Exception as e: return jsonify({'error': f'Write error: {str(e)}'}), 500 @app.route('/api/logs') def get_logs(): global log_buffer with log_lock: return jsonify({'logs': log_buffer}) @app.route('/api/preview') def list_previews(): # Scan for any images in the workspace/model directory model_dir = os.path.join(os.path.dirname(__file__), 'workspace', 'model') if not os.path.exists(model_dir): return jsonify([]) previews = [] try: for root, dirs, files in os.walk(model_dir): for f in files: if f.lower().endswith(('.jpg', '.png')): path = os.path.join(root, f) mtime = os.path.getmtime(path) rel_path = os.path.relpath(path, model_dir) # Exclude huge files if any size = os.path.getsize(path) if size < 5 * 1024 * 1024: # Less than 5MB previews.append({ 'name': f, 'path': rel_path, 'mtime': mtime, 'size': size }) # Sort by modification time (most recent first) previews.sort(key=lambda x: x['mtime'], reverse=True) except Exception as e: return jsonify({'error': str(e)}), 500 return jsonify(previews) @app.route('/api/preview/file/') def serve_preview_file(filename): model_dir = os.path.join(os.path.dirname(__file__), 'workspace', 'model') return send_from_directory(model_dir, filename) @app.route('/api/upload', methods=['POST']) def upload_file(): if 'file' not in request.files: return jsonify({'error': 'No file part'}), 400 file = request.files['file'] target = request.form.get('target') # 'src' or 'dst' if file.filename == '': return jsonify({'error': 'No selected file'}), 400 if target not in ('src', 'dst'): return jsonify({'error': 'Invalid upload target'}), 400 ext = os.path.splitext(file.filename)[1].lower() if ext not in ('.mp4', '.avi', '.mkv', '.mov'): return jsonify({'error': 'Unsupported file format'}), 400 workspace_dir = os.path.join(os.path.dirname(__file__), 'workspace') os.makedirs(workspace_dir, exist_ok=True) # Save as data_src. or data_dst. filename = f"data_{target}{ext}" dest_path = os.path.join(workspace_dir, filename) # Remove existing ones with other extensions if necessary, to keep DFL working cleanly for existing_ext in ('.mp4', '.avi', '.mkv', '.mov'): try: os.remove(os.path.join(workspace_dir, f"data_{target}{existing_ext}")) except FileNotFoundError: pass try: file.save(dest_path) return jsonify({'status': 'uploaded', 'filename': filename}) except Exception as e: return jsonify({'error': f'Failed to save file: {str(e)}'}), 500 if __name__ == '__main__': # Make sure workspace directories exist os.makedirs('templates', exist_ok=True) os.makedirs('static', exist_ok=True) port = 8082 print(f"Starting DeepFaceLab Web UI Server on port {port}...") print(f"Access it at http://localhost:{port}/") app.run(host='0.0.0.0', port=port, debug=False)