Files

445 lines
18 KiB
Python

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/<path:filename>')
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.<ext> or data_dst.<ext>
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)