diff --git a/roboflow/cli/handlers/train.py b/roboflow/cli/handlers/train.py index f2a69651..f597cb6a 100644 --- a/roboflow/cli/handlers/train.py +++ b/roboflow/cli/handlers/train.py @@ -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. @@ -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 @@ -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 @@ -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: diff --git a/tests/cli/test_train_handler.py b/tests/cli/test_train_handler.py index fcb1517b..2c35a8e1 100644 --- a/tests/cli/test_train_handler.py +++ b/tests/cli/test_train_handler.py @@ -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() @@ -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."""