File size: 9,078 Bytes
7c275a5 4607230 7c275a5 7d99559 2618db6 7c275a5 f9f499d 7c275a5 f9f499d 7c275a5 0193004 21dc5ba 7c275a5 4607230 7c275a5 4607230 7d99559 4607230 7c275a5 f9f499d 7c275a5 f9f499d 7c275a5 f9f499d 7c275a5 21dc5ba 7c275a5 21dc5ba 7c275a5 21dc5ba 7c275a5 | 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 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 | import os
import sys
import cv2
import torch
torch.set_num_threads(8) # Ryzen 7735HS has 8C/16T
# oneDNN can be unstable on the Ryzen CPU during BasicSR training.
torch.backends.mkldnn.enabled = False
import yaml
import subprocess
import glob
import argparse
def main():
parser = argparse.ArgumentParser(description="Train CodeFormer with custom parameters.")
parser.add_argument("--verify", action="store_true", help="Run 2 iterations for verification purposes.")
args = parser.parse_args()
project_dir = os.path.dirname(os.path.abspath(__file__))
codeformer_dir = os.path.join(project_dir, "models", "CodeFormer")
print("=== Custom CodeFormer Training Runner ===")
# 1. Check GPU availability
device = "cuda" if torch.cuda.is_available() else "cpu"
num_gpus = torch.cuda.device_count() if device == "cuda" else 0
print(f"Device detected: {device.upper()}")
print(f"Number of GPUs available: {num_gpus}")
# 2. Check and prepare dataset
dataset_dir = os.path.join(codeformer_dir, "datasets", "ffhq", "ffhq_512")
if not os.path.exists(dataset_dir) or len(os.listdir(dataset_dir)) == 0:
print("Dataset directory is empty. Preparing dataset images...")
try:
import prepare_toy_training
prepare_toy_training.main()
print("Dataset preparation completed.")
except Exception as e:
print(f"Error preparing dataset: {e}")
sys.exit(1)
else:
print(f"Dataset found at {dataset_dir} ({len(os.listdir(dataset_dir))} images).")
# 3. Update configuration file
config_path = os.path.join(codeformer_dir, "options", "CodeFormer_stage3_custom.yml")
if not os.path.exists(config_path):
print(f"Error: Config file not found at {config_path}")
sys.exit(1)
print(f"Reading configuration from {config_path}...")
with open(config_path, "r", encoding="utf-8") as f:
config = yaml.safe_load(f)
# Keep the deterministic benchmark holdout out of all training phases.
# It is intentionally ignored by Git because it may identify private data;
# create it with tools/prepare_benchmark.py before training.
holdout_manifest = os.path.join(project_dir, "benchmarks", "holdout_paths.txt")
if os.path.isfile(holdout_manifest):
for dataset in config.get("datasets", {}).values():
dataset["exclude_manifest"] = holdout_manifest
print(f"Excluding benchmark holdout from training: {holdout_manifest}")
else:
print("WARNING: benchmark holdout manifest not found; refusing quality claims for this run.")
# Save original total_iter BEFORE any runtime overrides — we must restore
# this after writing so --verify mode never permanently corrupts the file.
original_total_iter = config.get("train", {}).get("total_iter", 20000)
# Dynamically set GPU count
config["num_gpu"] = num_gpus
config["dist"] = False
config["dist_params"] = None
# Use multiple workers to saturate CPU cores when training on GPU; set to 0 on CPU to prevent Windows multiprocessing MemoryError/segfaults.
if "datasets" in config:
for phase in config["datasets"]:
dataset = config["datasets"][phase]
if device == "cpu":
dataset["num_worker_per_gpu"] = 0
# MUST be null on CPU — 'cpu' prefetch spawns multiprocessing
# workers which cause MemoryError/segfaults on Windows (Task 8).
dataset["prefetch_mode"] = None
else:
dataset["num_worker_per_gpu"] = 4
# Update weights path if they are in the project weights folder
project_weights_path = os.path.join(project_dir, "weights", "CodeFormer", "codeformer.pth")
if os.path.exists(project_weights_path):
# basicSR is relative to the running dir which is models/CodeFormer
config["path"]["pretrain_network_g"] = "../../weights/CodeFormer/codeformer.pth"
print(f"Configured pretrain generator path to: {config['path']['pretrain_network_g']}")
# 4. Auto-detect resume state from the latest experiment checkpoint
experiments_dir = os.path.join(codeformer_dir, "experiments")
latest_state = None
latest_state_iter = 0
# Search for existing experiment directories matching our config name pattern
exp_pattern = os.path.join(experiments_dir, "*_CodeFormer_stage3_custom", "training_states", "*.state")
state_files = glob.glob(exp_pattern)
for state_file in state_files:
basename = os.path.basename(state_file)
try:
iter_num = int(basename.replace(".state", ""))
if iter_num > latest_state_iter:
latest_state_iter = iter_num
latest_state = state_file
except ValueError:
continue
if latest_state:
print(f"\n>>> RESUME MODE: Found checkpoint at iteration {latest_state_iter}")
print(f" State file: {latest_state}")
config["path"]["resume_state"] = latest_state
config["path"]["pretrain_network_g"] = None
if args.verify:
config["train"]["total_iter"] = latest_state_iter + 2
print(f" Resuming training from iter {latest_state_iter} to {config['train']['total_iter']} (verification mode)...")
else:
if latest_state_iter >= config.get("train", {}).get("total_iter", 20000):
print(f"\n>>> Training already completed ({latest_state_iter} >= {config.get('train', {}).get('total_iter', 20000)} total_iter)")
sys.exit(0)
print(f" Resuming training from iter {latest_state_iter} to {config['train']['total_iter']}...")
else:
print("\n>>> FRESH START: No previous checkpoint found. Starting from pretrained weights.")
if args.verify:
config["train"]["total_iter"] = 2
print(f" Training from scratch to {config['train']['total_iter']} (verification mode)...")
else:
print(f" Training from scratch to {config.get('train', {}).get('total_iter', 20000)}...")
# 5. Write runtime config to a TEMP yml file — NEVER modify the original.
#
# Strategy: write all runtime overrides (num_gpu, prefetch_mode, resume_state,
# total_iter) to a separate temp file. basicsr/train.py reads from that temp
# file. The canonical yml with all comments is never touched again.
#
# This fixes two bugs permanently:
# (a) --verify mode used to persist total_iter=checkpoint+2 to disc, causing
# the next full run to stop after only 2 iterations.
# (b) yaml.dump was stripping all comments from the original yml every run.
import tempfile, shutil
tmp_fd, tmp_config_path = tempfile.mkstemp(suffix=".yml", prefix="cf_runtime_",
dir=os.path.join(codeformer_dir, "options"))
try:
with os.fdopen(tmp_fd, "w", encoding="utf-8") as f:
yaml.dump(config, f, default_flow_style=False, sort_keys=False)
print(f"Runtime config written to temp file: {os.path.basename(tmp_config_path)}")
print(f" device={device.upper()}, num_gpu={num_gpus}, total_iter={config['train']['total_iter']}")
# 6. Run the training process pointing at the temp config
train_script = os.path.join("basicsr", "train.py")
cmd = [
sys.executable,
train_script,
"-opt",
os.path.relpath(tmp_config_path, codeformer_dir),
"--launcher",
"none"
]
# Environment setup — CPU stability settings (Task 8 findings)
env = os.environ.copy()
for _k in ["OMP_NUM_THREADS", "MKL_NUM_THREADS", "OPENBLAS_NUM_THREADS",
"NUMEXPR_NUM_THREADS", "VECLIB_MAXIMUM_THREADS"]:
env[_k] = "8" # Match torch.set_num_threads — Ryzen 7735HS 8C/16T
env["OPENCV_OPENCL_RUNTIME"] = "disabled"
env["OPENCV_THREAD_LIMIT"] = "1"
cv2.setNumThreads(0)
env["PYTHONPATH"] = os.path.pathsep.join([codeformer_dir, env.get("PYTHONPATH", "")])
print("\nStarting training process. Command:")
print(" ".join(cmd))
print(f"Working directory: {codeformer_dir}")
print("------------------------------------------")
try:
result = subprocess.run(cmd, cwd=codeformer_dir, env=env, check=True)
print("------------------------------------------")
print("Training execution completed successfully!")
except subprocess.CalledProcessError as e:
print("------------------------------------------")
print(f"Training failed with exit code: {e.returncode}")
sys.exit(e.returncode)
finally:
# Always clean up the temp file after the run (success or failure)
if os.path.exists(tmp_config_path):
os.remove(tmp_config_path)
if __name__ == "__main__":
main()
|