Akegarasu / Akegarasu/lora-scripts

希望可以加入任务ID作为TensorBoard日志的前缀,这样可以通过接口获取到任务对应的loss值

Abierto
#695 0 comentarios 0 reacciones 0 asignados Ver en GitHub
Lenguaje dominante
Python
Estrellas
6.1k
Forks
699
Métricas de merge de PR
Sin PR fusionados en 30 d

Descripción

![Image](https://github.com/user-attachments/assets/bf8c16a6-d246-4d10-9881-b59bfaa0876c)
我正在基于秋叶训练器开发一款GUI训练管理工具,但发现可用api比较少,很多功能无法实现。例如获取任务的日志或者训练进度等信息,希望可以完善一下api能力。当前有一个硬性问题是无法获取任务对应的TensorBoard日志,保存的日志未和任务产生关联,因此通过查看源码发现在run_train方法中可以传入日志前缀,这样改动后即可在业务侧调用接口获取任务对应的日志数据,由此我获得到了日志的loss和训练step相关数据

```python
def run_train(toml_path: str,
trainer_file: str = "./scripts/train_network.py",
gpu_ids: Optional[list] = None,
cpu_threads: Optional[int] = 2):
log.info(f"Training started with config file / 训练开始,使用配置文件: {toml_path}")

customize_env = os.environ.copy()
customize_env["ACCELERATE_DISABLE_RICH"] = "1"
customize_env["PYTHONUNBUFFERED"] = "1"
customize_env["PYTHONWARNINGS"] = "ignore::FutureWarning,ignore::UserWarning"
# 创建任务ID
if not (task := tm.create_task([], customize_env)):
return APIResponse(status="error", message="Failed to create task / 无法创建训练任务")

args = [
sys.executable, "-m", "accelerate.commands.launch", # use -m to avoid python script executable error
"--num_cpu_threads_per_process", str(cpu_threads), # cpu threads
"--quiet", # silence accelerate error message
trainer_file,
"--config_file", toml_path,
"--log_prefix", f"{task.task_id}_",
]

if gpu_ids:
customize_env["CUDA_VISIBLE_DEVICES"] = ",".join(gpu_ids)
log.info(f"Using GPU(s) / 使用 GPU: {gpu_ids}")

if len(gpu_ids) > 1:
args[3:3] = ["--multi_gpu", "--num_processes", str(len(gpu_ids))]
if sys.platform == "win32":
customize_env["USE_LIBUV"] = "0"

task.command = args

def _run():
try:
task.execute()
result = task.communicate()
if result.returncode != 0:
log.error(f"Training failed / 训练失败")
else:
log.info(f"Training finished / 训练完成")
except Exception as e:
log.error(f"An error occurred when training / 训练出现致命错误: {e}")

coro = asyncio.to_thread(_run)
asyncio.create_task(coro)

return APIResponse(status="success", message=f"Training started / 训练开始 ID: {task.task_id}")
```

Guía de contribución

No hay ninguna guía de contribución indexada para este repositorio

Evaluación

Este issue todavía no se ha evaluado.

Recibe los nuevos issues en tu correo

Un resumen breve de issues de GitHub para principiantes.