Skip to content
Merged
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
29 changes: 14 additions & 15 deletions roboflow/cli/handlers/train.py
Original file line number Diff line number Diff line change
Expand Up @@ -88,8 +88,8 @@ def start_training(
) -> None:
"""Start training for a dataset version.

With --train-recipe, the training is created via the v2 trainings API
and the new trainingId is printed. Start from the ``template`` field of
With --train-recipe, the training uses v2 and prints its trainingId.
Action Recognition also uses v2 without a recipe. Start from the ``template`` field of
``roboflow train recipe`` output, edit it (hyperparameters, online
augmentation), and pass it inline or as ``@path/to/file.json``; --epochs is folded into its
hyperparameters unless the recipe already sets epochs.
Expand Down Expand Up @@ -297,11 +297,10 @@ def _start(args): # noqa: ANN001
output_error(args, "No API key found.", hint="Set ROBOFLOW_API_KEY or run 'roboflow auth login'.", exit_code=2)
return

# Custom recipes go through the v2 trainings API. Presence, not
# truthiness: an explicitly supplied empty value (e.g. an unset shell
# variable) must fail JSON validation, not fall through and start a
# legacy training.
if getattr(args, "train_recipe", None) is not None:
# Exact Cosmos 3 Edge and custom recipes use v2 so callers receive the
# trainingId. Presence, not truthiness: an explicitly supplied empty
# recipe must fail validation rather than start another training.
if args.model_type == "cosmos3-edge" or getattr(args, "train_recipe", None) is not None:
_start_v2(args, api_key, workspace_url, project_slug)
return

Expand Down Expand Up @@ -383,7 +382,7 @@ def _parse_json_flag(args, raw, flag):


def _start_v2(args, api_key, workspace_url, project_slug):
"""Create a training via the v2 trainings API with a custom trainRecipe."""
"""Create a v2 training, with a trainRecipe only when supplied."""
from roboflow.adapters import rfapi
from roboflow.cli._output import output, output_error
from roboflow.util.train_recipe import fold_epochs_into_recipe
Expand All @@ -400,13 +399,13 @@ def _start_v2(args, api_key, workspace_url, project_slug):
),
)
return
train_recipe = _parse_json_flag(args, args.train_recipe, "--train-recipe")
if args.epochs is not None:
# Fold --epochs into the recipe: the server dense-fills recipe
# hyperparameters (including a default epochs) and resolves them
# ahead of the body's top-level value, which would otherwise be
# silently ignored. An epochs set in the recipe wins.
train_recipe = fold_epochs_into_recipe(train_recipe, args.epochs)
train_recipe = args.train_recipe
if train_recipe is not None:
train_recipe = _parse_json_flag(args, train_recipe, "--train-recipe")
if args.epochs is not None:
# The server resolves recipe epochs ahead of the top-level value.
# An explicit recipe value wins over --epochs.
train_recipe = fold_epochs_into_recipe(train_recipe, args.epochs)

# Ensure the version has the required export format before training
if args.model_type:
Expand Down
95 changes: 95 additions & 0 deletions tests/cli/test_train_handler.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,10 +8,13 @@
import types
import unittest
from unittest.mock import MagicMock, patch
from urllib.parse import parse_qs, urlsplit

import responses
from typer.testing import CliRunner

from roboflow.cli import app
from roboflow.config import API_URL

runner = CliRunner()

Expand Down Expand Up @@ -464,9 +467,101 @@ def test_start_help_shows_train_recipe_flag_only(self) -> None:
self.assertEqual(result.exit_code, 0)
output = _strip_ansi(result.output)
self.assertIn("--train-recipe", output)
self.assertIn("Action Recognition", output)
self.assertIn("trainingId", output)
self.assertNotIn("--hyperparameters", output)


class TestActionRecognitionTrainCLI(unittest.TestCase):
"""Exercise the public train command through its real HTTP adapters."""

BASE = f"{API_URL}/audit-ws/audit-project/1"

def _invoke(self, model_type, *, json_output=True, recipe=None):
args = ["--api-key", "fixture-key", "--workspace", "audit-ws", "--quiet"]
if json_output:
args.append("--json")
args.extend(["train", "start", "-p", "audit-project", "-v", "1", "-t", model_type, "--epochs", "5"])
if recipe is not None:
args.extend(["--train-recipe", recipe])
return runner.invoke(app, args)

def test_recipe_free_cosmos_exports_video_coco_and_returns_id_in_json_and_text(self):
for json_output in (True, False):
with self.subTest(json_output=json_output), responses.RequestsMock() as mocked, patch("time.sleep"):
mocked.add(responses.GET, self.BASE, json={"version": {"generating": False, "exports": []}})
mocked.add(responses.GET, f"{self.BASE}/video-coco", json={"export": {"link": "https://fixture.test"}})
mocked.add(responses.GET, self.BASE, json={"version": {"generating": False, "exports": ["video-coco"]}})
mocked.add(
responses.POST,
f"{self.BASE}/v2/trainings",
json={"trainingId": "audit-t1", "status": "queued", "jobId": "audit-j1"},
)

result = self._invoke("cosmos3-edge", json_output=json_output)

self.assertEqual(result.exit_code, 0, result.output)
self.assertIn("audit-t1", result.stdout)
if json_output:
self.assertEqual(json.loads(result.stdout)["trainingId"], "audit-t1")
paths = [urlsplit(call.request.url).path for call in mocked.calls]
self.assertEqual(
paths,
[
"/audit-ws/audit-project/1",
"/audit-ws/audit-project/1/video-coco",
"/audit-ws/audit-project/1",
"/audit-ws/audit-project/1/v2/trainings",
],
)
create = mocked.calls[-1].request
self.assertEqual(parse_qs(urlsplit(create.url).query)["api_key"], ["fixture-key"])
self.assertEqual(json.loads(create.body), {"modelType": "cosmos3-edge", "epochs": 5})

def test_recipe_free_other_models_keep_legacy_route(self):
for model_type in ("rfdetr-medium", "cosmos3-edge-vlm"):
with self.subTest(model_type=model_type), responses.RequestsMock() as mocked:
mocked.add(
responses.GET,
self.BASE,
json={"version": {"generating": False, "exports": ["coco", "yolov5pytorch"]}},
)
mocked.add(responses.POST, f"{self.BASE}/train", status=204)

result = self._invoke(model_type)

self.assertEqual(result.exit_code, 0, result.output)
self.assertEqual(json.loads(result.stdout)["status"], "training_started")
self.assertEqual(
[urlsplit(call.request.url).path for call in mocked.calls][-1], "/audit-ws/audit-project/1/train"
)
self.assertEqual(json.loads(mocked.calls[-1].request.body)["modelType"], model_type)

def test_explicit_recipe_still_folds_epochs_and_uses_v2(self):
with responses.RequestsMock() as mocked:
mocked.add(responses.GET, self.BASE, json={"version": {"generating": False, "exports": ["video-coco"]}})
mocked.add(responses.POST, f"{self.BASE}/v2/trainings", json={"trainingId": "audit-t2", "status": "queued"})

result = self._invoke("cosmos3-edge", recipe='{"schema_version":1,"hyperparameters":{"lr":0.1}}')

self.assertEqual(result.exit_code, 0, result.output)
self.assertEqual(json.loads(result.stdout)["trainingId"], "audit-t2")
self.assertEqual(
json.loads(mocked.calls[-1].request.body)["trainRecipe"]["hyperparameters"],
{"lr": 0.1, "epochs": 5},
)
self.assertEqual(len(mocked.calls), 2)

def test_malformed_explicit_recipe_does_not_start_training(self):
for recipe in ("", "{invalid"):
with self.subTest(recipe=recipe), responses.RequestsMock() as mocked:
result = self._invoke("cosmos3-edge", recipe=recipe)

self.assertEqual(result.exit_code, 1)
self.assertIn("Invalid JSON", result.output)
self.assertEqual(len(mocked.calls), 0)


class TestTrainSubcommandsRegister(unittest.TestCase):
"""train cancel/stop/results subcommands register correctly."""

Expand Down
Loading