ReVID / multinode_rl.py
GuoruiSong's picture
Add files using upload-large-folder tool
3e90852 verified
Raw
History Blame Contribute Delete
14.3 kB
import os
import subprocess
import shlex
from omegaconf import DictConfig, ListConfig, OmegaConf
def run_local(cmd: str, check: bool = True) -> None:
subprocess.run(f"bash -lc {shlex.quote(cmd)}", shell=True, check=check)
def run_local_async(cmd: str) -> subprocess.Popen:
return subprocess.Popen(f"bash -lc {shlex.quote(cmd)}", shell=True)
def run_remote(host: str, cmd: str, check: bool = True) -> None:
ssh_cmd = f'ssh root@{host} "bash -lc {shlex.quote(cmd)}"'
subprocess.run(ssh_cmd, shell=True, check=check)
def run_remote_async(host: str, cmd: str) -> subprocess.Popen:
ssh_cmd = f'ssh root@{host} "bash -lc {shlex.quote(cmd)}"'
return subprocess.Popen(ssh_cmd, shell=True)
def get_config():
cli_conf = OmegaConf.from_cli()
yaml_conf = OmegaConf.load(cli_conf.config)
return OmegaConf.merge(yaml_conf, cli_conf)
def begin_with(file_name: str):
with open(file_name, "w"):
pass
def make_init_bash(cfg) -> str:
sc = cfg.system
http_proxy = sc.HTTP_PROXY
https_proxy = sc.HTTP_PROXY
hf_home = sc.HF_HOME
envs_dir = sc.envs_dir
lines = []
lines.append("set -e")
if http_proxy is not None:
lines.append(f"echo 'export HTTP_PROXY={http_proxy}' >> ~/.bashrc")
if https_proxy is not None:
lines.append(f"echo 'export HTTPS_PROXY={https_proxy}' >> ~/.bashrc")
if hf_home is not None:
lines.append(f"echo 'export HF_HOME={hf_home}' >> ~/.bashrc")
lines.append("")
if envs_dir is not None:
lines.append(f"conda config --append envs_dirs {envs_dir} || true")
lines.append("")
lines.append("echo 'source ~/.bashrc' >> ~/.bash_profile")
lines.append("")
return "\n".join(lines)
if __name__ == "__main__":
def init_node(host: str):
run_remote(host, INIT_BASH, check=False)
def init_hosts(worker_hosts):
for h in worker_hosts:
init_node(h)
def env_prefix() -> str:
return (
"source ~/.bashrc && "
f"source activate {env_name} && "
)
def sample(worker_hosts, epoch, cfg, type, model_base, block_size = None, top_k = None, remasking_strategy = None):
project = cfg.experiment.project
procs = []
if model_base == "dream":
script_name = "dream_rl_rollout.py"
elif model_base == "llada" or model_base == "mmada":
script_name = "llada_rl_rollout.py"
elif model_base == "sdar":
script_name = "sdar_rl_rollout.py"
elif model_base == "trado":
script_name = "trado_rl_rollout.py"
for idx, host in enumerate(worker_hosts):
body = (
f"cd {BASE_DIR}/sample && "
f"python {script_name} "
f"config=../configs/{project}.yaml "
f"experiment.current_epoch={epoch} "
f"experiment.function={type} "
f"evaluation.top_k={top_k} "
f"evaluation.remasking_strategy={remasking_strategy} "
f"evaluation.block_size={block_size} "
f"experiment.node_index={idx}"
)
full_cmd = env_prefix() + body
if idx == 0:
procs.append(run_local_async(full_cmd))
else:
procs.append(run_remote_async(host, full_cmd))
for p in procs:
p.wait()
def execute(worker_hosts, epoch, cfg, type):
project = cfg.experiment.project
procs = []
for idx, host in enumerate(worker_hosts):
full_cmd = env_prefix() + (
f"cd {BASE_DIR}/reward && "
f"python rl_execute.py "
f"config=../configs/{project}.yaml "
f"experiment.function={type} "
f"experiment.current_epoch={epoch} "
f"experiment.node_index={idx}"
)
if idx == 0:
procs.append(run_local_async(full_cmd))
else:
procs.append(run_remote_async(host, full_cmd))
for p in procs:
p.wait()
def aggregate(epoch, cfg, type):
project = cfg.experiment.project
full_cmd = env_prefix() + (
f"cd {BASE_DIR}/reward && "
f"python rl_aggregate_data.py "
f"config=../configs/{project}.yaml "
f"experiment.function={type} "
f"experiment.current_epoch={epoch}"
)
run_local(full_cmd)
def reward(epoch, cfg, type, is_code_task, block_size = None, top_k = None, remasking_strategy = None):
project = cfg.experiment.project
if is_code_task:
script_name = "rl_code_reward.py"
else:
script_name = "rl_reward.py"
full_cmd = env_prefix() + (
f"cd {BASE_DIR}/reward && "
f"python {script_name} "
f"config=../configs/{project}.yaml "
f"experiment.function={type} "
f"evaluation.block_size={block_size} "
f"evaluation.top_k={top_k} "
f"evaluation.remasking_strategy={remasking_strategy} "
f"experiment.current_epoch={epoch}"
)
run_local(full_cmd)
def train(worker_hosts, epoch, cfg, model_base, target = None):
project = cfg.experiment.project
ds_file = cfg.experiment.deepspeed_file
num_nodes = len(worker_hosts)
master_ip = os.environ["MLP_WORKER_0_HOST"]
master_port = os.environ["MLP_WORKER_0_PORT"]
if target is None:
if model_base == "dream":
script_name = "rl_dream.py"
elif model_base == "llada":
script_name = "rl_llada.py"
elif model_base == "mmada":
script_name = "rl_mmada.py"
elif model_base == "sdar":
script_name = "rl_sdar.py"
elif model_base == "trado":
script_name = "rl_trado.py"
elif target == "policy":
if model_base == "sdar":
script_name = "train_sdar_policy.py"
elif model_base == "trado":
script_name = "train_trado_policy.py"
elif target == "value":
if model_base == "sdar":
script_name = "train_sdar_value.py"
elif model_base == "trado":
script_name = "train_trado_value.py"
procs = []
for idx, host in enumerate(worker_hosts):
body = (
f"cd {BASE_DIR} && "
"export DS_SKIP_CUDA_CHECK=1 && "
"accelerate launch "
f"--num_machines {num_nodes} "
f"--machine_rank {idx} "
f"--main_process_ip {master_ip} "
f"--main_process_port {master_port} "
f"--config_file accelerate_configs/{ds_file}.yaml "
f"train/{script_name} "
f"config=configs/{project}.yaml "
f"experiment.current_epoch={epoch}"
)
full_cmd = env_prefix() + body
if idx == 0:
procs.append(run_local_async(full_cmd))
else:
procs.append(run_remote_async(host, full_cmd))
print(f"[DISPATCH] train node {idx}{host}")
for p in procs:
p.wait()
print("All train nodes finished.")
def process_reward(worker_hosts, epoch, cfg):
project = cfg.experiment.project
# --- Step A: node0 preprocess ---
cmd_div = (
env_prefix() +
f"cd {BASE_DIR}/reward && "
f"python rl_process_divide_data.py "
f"config=../configs/{project}.yaml "
f"experiment.current_epoch={epoch}"
)
run_local(cmd_div)
print("[OK] rl_process_divide_data.py (node0) done.")
# --- Step B: multinode LLM rewarding ---
procs = []
for idx, host in enumerate(worker_hosts):
cmd_llm = (
env_prefix() +
f"cd {BASE_DIR}/sample && "
f"python llm_process_reward.py "
f"config=../configs/{project}.yaml "
f"experiment.current_epoch={epoch} "
f"experiment.node_index={idx}"
)
if idx == 0:
procs.append(run_local_async(cmd_llm))
else:
procs.append(run_remote_async(host, cmd_llm))
print(f"[DISPATCH] llm_process_reward node {idx}{host}")
for p in procs:
p.wait()
print("[OK] llm_process_reward.py (all nodes) done.")
# --- Step C: node0 aggregate ---
aggregate(epoch, cfg, "train")
print("[OK] aggregation done.")
# --- Step D: node0 final process ---
cmd_fin = (
env_prefix() +
f"cd {BASE_DIR}/reward && "
f"python rl_process_reward.py "
f"config=../configs/{project}.yaml "
f"experiment.current_epoch={epoch}"
)
run_local(cmd_fin)
print("[OK] rl_process_reward.py (node0) done.")
def init_value_model(epoch, cfg):
project = cfg.experiment.project
full_cmd = env_prefix() + (
f"cd {BASE_DIR}/train && "
f"python init_sdar_value_model.py "
f"config=../configs/{project}.yaml "
f"experiment.current_epoch={epoch}"
)
run_local(full_cmd)
cfg = get_config()
INIT_BASH = make_init_bash(cfg)
BASE_DIR = cfg.system.base_dir
env_name = cfg.system.env_name
total_step = cfg.experiment.total_step
project = cfg.experiment.project
num_node = cfg.experiment.num_node
worker_hosts = [os.environ[f"MLP_WORKER_{i}_HOST"] for i in range(num_node)]
from omegaconf import MISSING
if OmegaConf.select(cfg, "model.value_base_model", default=MISSING) is not MISSING:
have_value_model = True
else:
have_value_model = False
import time
time.sleep(30)
init_hosts(worker_hosts)
import time
time.sleep(10)
if cfg.experiment.start_from_scratch:
os.makedirs(f"{project}/results", exist_ok=True)
optimized = f"../{project}/ckpt/{cfg.model.optimized_name}"
path = (
f"{project}/results/results-rl-"
f"{optimized.replace('/', '.')}-"
f"{cfg.dataset.train_dataset}.txt"
)
begin_with(path)
path = (
f"{project}/results/results-eval-"
f"{optimized.replace('/', '.')}-"
f"{cfg.evaluation.eval_dataset}.txt"
)
begin_with(path)
if have_value_model:
init_value_model(1, cfg)
optimized_value = f"../{project}/ckpt/{cfg.model.optimized_value_name}"
path = (
f"{project}/results/results-rl-"
f"{optimized_value.replace('/', '.')}-"
f"{cfg.dataset.train_dataset}.txt"
)
begin_with(path)
epoch = cfg.experiment.current_epoch
model_base = cfg.model.model_base
if cfg.dataset.data_type == "code":
is_code_task = True
else:
is_code_task = False
if OmegaConf.select(cfg, "model.process_reward_model", default=MISSING) is not MISSING and cfg.model.process_reward_model is not None:
is_process_reward = True
else:
is_process_reward = False
while epoch <= total_step:
print(f"\n========== epoch {epoch} ==========")
sample(worker_hosts, epoch, cfg, "train", model_base)
if is_code_task:
execute(worker_hosts, epoch, cfg, "train")
aggregate(epoch, cfg, "train")
if is_process_reward:
process_reward(worker_hosts, epoch, cfg)
else:
reward(epoch, cfg, "train", is_code_task)
if have_value_model:
train(worker_hosts, epoch, cfg, model_base, target = "value")
train(worker_hosts, epoch, cfg, model_base, target = "policy")
else:
train(worker_hosts, epoch, cfg, model_base, target = None)
if epoch % cfg.experiment.eval_every == 0:
if model_base in ["sdar", "trado"]:
remasking_strategy_list = cfg.evaluation.remasking_strategy
top_k_list = cfg.evaluation.top_k
block_size = cfg.evaluation.block_size
for j in range(len(remasking_strategy_list)):
remasking_strategy = remasking_strategy_list[j]
top_k = top_k_list[j]
sample(worker_hosts, epoch, cfg, "evaluation", model_base, block_size = block_size, top_k = top_k, remasking_strategy = remasking_strategy)
if is_code_task:
execute(worker_hosts, epoch, cfg, "evaluation")
aggregate(epoch, cfg, "evaluation")
reward(epoch, cfg, "evaluation", is_code_task, block_size = block_size, top_k = top_k, remasking_strategy = remasking_strategy)
else:
block_size_list = cfg.evaluation.block_size
remasking_strategy_list = cfg.evaluation.remasking_strategy
if OmegaConf.select(cfg, "evaluation.top_k", default=MISSING) is not MISSING:
top_k = cfg.evaluation.top_k
else:
top_k = None
for j in range(len(remasking_strategy_list)):
remasking_strategy = remasking_strategy_list[j]
if model_base == "dream":
block_size = block_size_list[j]
elif model_base == "llada" or model_base == "mmada":
block_size = cfg.evaluation.block_size
sample(worker_hosts, epoch, cfg, "evaluation", model_base, block_size = block_size, top_k = top_k, remasking_strategy = remasking_strategy)
if is_code_task:
execute(worker_hosts, epoch, cfg, "evaluation")
aggregate(epoch, cfg, "evaluation")
reward(epoch, cfg, "evaluation", is_code_task, block_size = block_size, top_k = top_k, remasking_strategy = remasking_strategy)
epoch += 1