Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
64 changes: 64 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -84,6 +84,70 @@ uvx nanoflow run experiments/mdn-ablation.toml
uvx nanoflow run experiments/makd-ablation.toml
```

如果只想快速验证训练链路是否正常,也可以直接限制训练轮数或每轮 step 数:

```bash
uv run emotion-recognize train \
configs/dataset/MELD--E.toml \
configs/encoders/T+A+V.toml \
configs/fusion/DF-1.0.toml \
configs/fusion/kwargs/attn.toml \
configs/losses/classification/weight.toml \
--num-epochs 1 \
--max-train-batches 50 \
--seed 42
```

#### 使用 Modal 运行低成本 smoke test

仓库新增了 `src/recognize_modal/app.py`,用于自动化:

- 构建训练用 Modal Image
- 预热 Hugging Face 权重缓存 Volume
- 将 `datasets/` 和 `checkpoints/` 直接挂载到远端仓库路径
- 以受控的 `num_epochs` / `max_train_batches` 运行 smoke train

本地只需要安装 Modal 相关依赖:

```bash
uv sync --extra modal --dev
```

创建训练所需的 Volumes:

```bash
uv run modal volume create emotion-recognition-hf-cache
uv run modal volume create emotion-recognition-datasets
uv run modal volume create emotion-recognition-checkpoints
```

其中 `emotion-recognition-datasets` 里的目录结构需要与本地 `datasets/` 保持一致,例如 `MELD/`、`IEMOCAP/` 等子目录。

可选:先预热默认三模态实验会用到的 Hugging Face 权重:

```bash
uv run modal run src/recognize_modal/app.py::warm_hf_cache
```

然后运行一轮默认的 MELD 三模态 smoke train:

```bash
uv run modal run src/recognize_modal/app.py::train_smoke \
--num-epochs 1 \
--max-train-batches 50 \
--seed 42
```

默认 smoke 配置对应:

- `configs/dataset/MELD--E.toml`
- `configs/encoders/T+A+V.toml`
- `configs/fusion/DF-1.0.toml`
- `configs/fusion/kwargs/attn.toml`
- `configs/losses/classification/weight.toml`

训练生成的检查点会落到 `emotion-recognition-checkpoints` Volume,对应远端仓库内的 `./checkpoints/` 路径。

## 核心特性

### 🚀 高效训练机制
Expand Down
6 changes: 5 additions & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,10 @@ emotion-recognize = "recognize_cli.cli_recognize:app"
emotion-tool = "recognize_cli.cli_tool:app"

[project.optional-dependencies]
modal = [
"huggingface-hub>=0.28.1",
"modal>=1.1.0",
]
train = [
"bitsandbytes>=0.44.1",
"opencv-python>=4.0.0",
Expand Down Expand Up @@ -82,7 +86,7 @@ ignore = ["F841", "PGH003", "B008", "RUF001"]
"__init__.py" = ["F401", "I002"]

[tool.ruff.lint.isort]
known-first-party = ["recognize"]
known-first-party = ["recognize", "recognize_modal"]
required-imports = ["from __future__ import annotations"]

[tool.ruff.lint.flake8-tidy-imports]
Expand Down
21 changes: 17 additions & 4 deletions src/recognize/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -135,11 +135,18 @@ def train_epoch(
train_data_loader: DataLoader,
*,
dropout_prob: float | None = None,
max_train_batches: int | None = None,
update_hook: Callable[[int, float], None] | None = None,
):
model.train()
trainer.clear_losses()
batch_num = len(train_data_loader)
if max_train_batches is not None:
batch_num = min(batch_num, max_train_batches)
log_interval = max(1, batch_num // 10)
for batch_index, batch in enumerate(train_data_loader):
if batch_index >= batch_num:
break
# NOTE: randomly remove every modality independently
if isinstance(batch, LazyMultimodalInput) and dropout_prob is not None:
if random.random() < dropout_prob:
Expand All @@ -148,7 +155,8 @@ def train_epoch(
batch.video_paths = None
trainer.train_batch(batch)
if update_hook is not None:
if batch_index % (len(train_data_loader) // 10):
processed_batches = batch_index + 1
if processed_batches % log_interval == 0 or processed_batches == batch_num:
loss_value = trainer.loss_mean()
trainer.clear_losses()
else:
Expand All @@ -169,12 +177,16 @@ def train_and_eval(
use_valid: bool = True,
eval_interval: int = 1,
dropout_prob: float | None = None,
max_train_batches: int | None = None,
):
batch_num = len(train_data_loader)
if max_train_batches is not None:
batch_num = min(batch_num, max_train_batches)
trainer = get_trainer(
model,
{"train": train_data_loader, "valid": valid_data_loader, "test": test_data_loader or valid_data_loader},
len(train_data_loader),
len(train_data_loader) * num_epochs,
batch_num,
batch_num * num_epochs,
)
checkpoint_dir.mkdir(parents=True, exist_ok=True)
if test_data_loader is None:
Expand All @@ -196,7 +208,6 @@ def train_and_eval(
stopper.history = [(epoch, history) for epoch, history in stopper.history if epoch <= epoch_start]

last_better_epoch = stopper.last_better_epoch
batch_num = len(train_data_loader)
result: TrainingResult | None = None

logger.info(f"Train model [blue]{model_label}[/]. Save to [blue]{checkpoint_dir}[/]")
Expand All @@ -220,10 +231,12 @@ def train_and_eval(
batch_index=0,
)
for epoch in range(epoch_start + 1, num_epochs):
progress.update(task, batch_index=0)
train_epoch(
model,
trainer,
train_data_loader,
max_train_batches=max_train_batches,
update_hook=lambda batch_index, loss_value: progress.update(
task, loss=loss_value, batch_index=batch_index + 1
),
Expand Down
11 changes: 10 additions & 1 deletion src/recognize_cli/cli_recognize.py
Original file line number Diff line number Diff line change
Expand Up @@ -178,6 +178,13 @@ def train(
seed: int | None = None,
checkpoint: Path | None = None,
from_checkpoint: Path | None = None,
num_epochs: int = typer.Option(200, min=1, help="Number of training epochs"),
max_train_batches: int | None = typer.Option(
None,
min=1,
help="Maximum number of training batches to process per epoch. Useful for smoke tests.",
),
eval_interval: int = typer.Option(1, min=1, help="Evaluate every N epochs"),
teacher_checkpoint: Path | None = typer.Option(
None, help="The checkpoint of the teacher model to be used in distillation"
),
Expand Down Expand Up @@ -342,10 +349,12 @@ def train(
dev_data_loader,
test_data_loader,
checkpoint_dir=checkpoint_dir,
num_epochs=200,
num_epochs=num_epochs,
model_label=model_label,
use_valid=False,
eval_interval=eval_interval,
dropout_prob=config.dropout_prob,
max_train_batches=max_train_batches,
)
logger.info(f"Test result in best model({model_label}):")
result.print()
Expand Down
1 change: 1 addition & 0 deletions src/recognize_modal/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
from __future__ import annotations
201 changes: 201 additions & 0 deletions src/recognize_modal/app.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,201 @@
from __future__ import annotations

import os
import re
import shlex
import tomllib
from collections.abc import Sequence
from pathlib import Path, PurePosixPath
from typing import Any

import modal

PROJECT_ROOT = Path(__file__).resolve().parents[2]
REMOTE_PROJECT_ROOT = Path("/root/emotion-recognition")
REMOTE_HF_CACHE_ROOT = Path("/vol/hf")
REMOTE_DATASETS_ROOT = REMOTE_PROJECT_ROOT / "datasets"
REMOTE_CHECKPOINTS_ROOT = REMOTE_PROJECT_ROOT / "checkpoints"

HF_CACHE_VOLUME_NAME = "emotion-recognition-hf-cache"
DATASETS_VOLUME_NAME = "emotion-recognition-datasets"
CHECKPOINTS_VOLUME_NAME = "emotion-recognition-checkpoints"

DEFAULT_MODEL_IDS = [
"roberta-large",
"facebook/data2vec-audio-base-960h",
"MCG-NJU/videomae-base-finetuned-kinetics",
]
DEFAULT_SMOKE_CONFIG_PATHS = [
"configs/dataset/MELD--E.toml",
"configs/encoders/T+A+V.toml",
"configs/fusion/DF-1.0.toml",
"configs/fusion/kwargs/attn.toml",
"configs/losses/classification/weight.toml",
]


def _dependency_name(specifier: str) -> str:
return re.split(r"[<>=!~;\\[ ]", specifier, maxsplit=1)[0]


def _load_modal_dependencies() -> tuple[list[str], str]:
with open(PROJECT_ROOT / "pyproject.toml", "rb") as file:
pyproject = tomllib.load(file)

project = pyproject["project"]
dependencies = [
*project["dependencies"],
*project["optional-dependencies"]["train"],
*project["optional-dependencies"]["modal"],
]

filtered_dependencies: list[str] = []
torch_dependency: str | None = None
seen_dependencies: set[str] = set()
for dependency in dependencies:
dependency_name = _dependency_name(dependency)
if dependency_name == "emotion-recognition-utils":
continue
if dependency_name == "torch":
torch_dependency = dependency
continue
if dependency_name in seen_dependencies:
continue
seen_dependencies.add(dependency_name)
filtered_dependencies.append(dependency)

if torch_dependency is None:
raise ValueError("torch dependency not found in pyproject.toml")
return filtered_dependencies, torch_dependency


def _recursive_update(base: dict[str, Any], update: dict[str, Any]) -> dict[str, Any]:
for key, value in update.items():
if isinstance(value, dict) and isinstance(base.get(key), dict):
base[key] = _recursive_update(base[key], value)
else:
base[key] = value
return base


def _load_training_config_dict(config_paths: Sequence[str]) -> dict[str, Any]:
config: dict[str, Any] = {}
for config_path in config_paths:
with open(PROJECT_ROOT / config_path, "rb") as file:
_recursive_update(config, tomllib.load(file))
return config


def _resolve_model_ids(config_paths: Sequence[str] | None, model_ids: Sequence[str] | None) -> list[str]:
resolved_model_ids = list(DEFAULT_MODEL_IDS)
selected_config_paths = list(config_paths or DEFAULT_SMOKE_CONFIG_PATHS)
if selected_config_paths:
config = _load_training_config_dict(selected_config_paths)
encoder_configs = config.get("model", {}).get("encoder", {})
for encoder_config in encoder_configs.values():
if isinstance(encoder_config, dict) and isinstance(encoder_config.get("model"), str):
resolved_model_ids.append(encoder_config["model"])
if model_ids is not None:
resolved_model_ids.extend(model_ids)
return list(dict.fromkeys(resolved_model_ids))


REMOTE_ENV = {
"HF_HOME": REMOTE_HF_CACHE_ROOT.as_posix(),
"HF_HUB_CACHE": (REMOTE_HF_CACHE_ROOT / "hub").as_posix(),
"PYTHONPATH": ":".join(
[
(REMOTE_PROJECT_ROOT / "src").as_posix(),
(REMOTE_PROJECT_ROOT / "packages" / "emotion-recognition-utils" / "src").as_posix(),
]
),
"TOKENIZERS_PARALLELISM": "false",
"TRANSFORMERS_CACHE": (REMOTE_HF_CACHE_ROOT / "hub").as_posix(),
}

IMAGE_DEPENDENCIES, TORCH_DEPENDENCY = _load_modal_dependencies()

hf_cache_volume = modal.Volume.from_name(HF_CACHE_VOLUME_NAME, create_if_missing=True)
datasets_volume = modal.Volume.from_name(DATASETS_VOLUME_NAME, create_if_missing=True)
checkpoints_volume = modal.Volume.from_name(CHECKPOINTS_VOLUME_NAME, create_if_missing=True)

image = (
modal.Image.debian_slim(python_version="3.12")
.apt_install("ffmpeg", "git", "libgl1", "libsndfile1")
.uv_pip_install(*IMAGE_DEPENDENCIES)
.run_commands(
"python -m pip install --no-cache-dir --index-url https://download.pytorch.org/whl/cu126 "
f"{shlex.quote(TORCH_DEPENDENCY)}"
)
.add_local_dir(PROJECT_ROOT.as_posix(), remote_path=REMOTE_PROJECT_ROOT.as_posix())
.env(REMOTE_ENV)
.workdir(REMOTE_PROJECT_ROOT.as_posix())
)

app = modal.App(name="emotion-recognition")

BASE_VOLUMES: dict[str | PurePosixPath, modal.Volume | modal.CloudBucketMount] = {
REMOTE_HF_CACHE_ROOT.as_posix(): hf_cache_volume,
REMOTE_DATASETS_ROOT.as_posix(): datasets_volume,
REMOTE_CHECKPOINTS_ROOT.as_posix(): checkpoints_volume,
}


@app.function(image=image, cpu=2.0, timeout=1800, volumes=BASE_VOLUMES)
def warm_hf_cache(
config_paths: list[str] | None = None,
model_ids: list[str] | None = None,
) -> list[str]:
from huggingface_hub import snapshot_download

resolved_model_ids = _resolve_model_ids(config_paths, model_ids)
cache_dir = REMOTE_HF_CACHE_ROOT / "hub"
for model_id in resolved_model_ids:
snapshot_download(
repo_id=model_id,
cache_dir=cache_dir,
ignore_patterns=["*.h5", "*.msgpack", "*.onnx"],
)

hf_cache_volume.commit()
return resolved_model_ids


@app.function(
image=image,
gpu="A10G",
cpu=8.0,
memory=32768,
timeout=3600,
volumes=BASE_VOLUMES,
)
def train_smoke(
config_paths: list[str] | None = None,
batch_size: int | None = None,
seed: int | None = None,
num_epochs: int = 1,
max_train_batches: int = 10,
eval_interval: int = 1,
checkpoint: str | None = None,
) -> str:
from recognize_cli.cli_recognize import train

os.chdir(REMOTE_PROJECT_ROOT)
resolved_config_paths = [Path(path) for path in config_paths or DEFAULT_SMOKE_CONFIG_PATHS]
checkpoint_path = Path(checkpoint) if checkpoint is not None else None

train(
resolved_config_paths,
batch_size=batch_size,
seed=seed,
checkpoint=checkpoint_path,
num_epochs=num_epochs,
max_train_batches=max_train_batches,
eval_interval=eval_interval,
)

hf_cache_volume.commit()
checkpoints_volume.commit()
if checkpoint_path is not None:
return checkpoint_path.as_posix()
return REMOTE_CHECKPOINTS_ROOT.as_posix()
Loading
Loading