Skip to content

Monitor

Monitor

Bases: Callback

Monitor callback for logging training metrics to MLflow.

This callback logs training progress and evaluation metrics to MLflow for experiment tracking and visualization.

What gets logged on a fit() run:

  • Metrics per epoch (step=epoch): every training metric, the val_* validation metrics, the LM/EM spend of the epoch (epoch_cost, epoch_tokens, per phase epoch_cost_inference|reward|optimizer), the running spend of the fit (fit_cost, fit_tokens), the all-time in-process spend (total_cost, total_tokens) and epoch_duration_s; fit_duration_s at the end. With log_batch_metrics=True, batch metrics are logged too, prefixed train_batch_ / val_batch_ with their own step counters.
  • Params: the fit arguments (epochs, batch/minibatch size, validation split/freq, dataset sizes), the compile configuration (reward, metrics, optimizer and its config), the program (name, class, number of modules and trainable variables) and every language or embedding model reachable from the program (lm.<name>.model, api base, sampling settings).
  • Datasets: the train and validation sets as MLflow dataset inputs, one row per sample with JSON inputs and expectations columns.
  • Model: the program (architecture and trained variables, the JSON of Program.save()) as an MLflow pyfunc model wrapped in SynalinksProgramModel, so it can be loaded with mlflow.pyfunc.load_model() and served with mlflow models serve. A new version is logged every time val_reward improves (reward when there is no validation data, the end of training when neither is available), registered in the Model Registry under the program's name; the best one is tagged synalinks.best. Each is an MLflow LoggedModel linked to the run; its id is written into the saved program (program._mlflow_model_id) so later evaluate() runs of the same program link their traces and metrics to that version and nest under the training run.
  • Prompts: each trainable Generator-like module's prompt (the system turn rendered from its optimized instructions and examples, plus a {{ inputs }} user turn, the module's output schema as response_format, and its LM settings as model_config) is registered in the MLflow Prompt Registry under <program>.<module>. A new version is created at every epoch where the rendered prompt changed, its commit message and val_reward tag recording that epoch's validation reward; the best version gets the best alias at the end. Versions are linked to the run, to the logged model, and to the traces of the module's later calls.
  • Artifacts: the program plot at the beginning of training.

A standalone evaluate() opens its own run (type genai_evaluate), logs the evaluation metrics, eval_cost, eval_tokens, eval_duration_s, the evaluated dataset, and every sample's reward as a reward feedback assessment on its trace.

Parameters:

Name Type Description Default
experiment_name str

Name of the MLflow experiment. If None, uses the experiment of synalinks.enable_observability() when it was called (so runs land next to their traces), else the program name.

None
run_name str

Name of the MLflow run. If None, the program's name, suffixed _train or _test.

None
tracking_uri str

MLflow tracking server URI. If None, uses the value from synalinks.enable_observability() or the default (local ./mlruns directory or MLFLOW_TRACKING_URI env var).

None
log_batch_metrics bool

Whether to log metrics at batch level (default: False).

False
log_epoch_metrics bool

Whether to log metrics at epoch level (default: True).

True
log_program_plot bool

Whether to log the program plot as an artifact at the beginning of training (default: True).

True
log_program_model bool

Whether to log the program as an MLflow model (default: True).

True
tags dict

Optional tags to add to the MLflow run.

None
run_id str

Optional. The id of an existing MLflow run to resume instead of starting a new one. Metrics keep being appended to it, the step counters continuing after the last step already logged.

None
resume bool

Whether to look up an existing run named run_name in the experiment and resume it (default: False). Creates the run the first time. This is how repeated evaluate() calls draw a chart over time: each evaluation adds one point to the metrics of the same run.

False
log_assessments bool

Whether to log each evaluated sample's reward as a reward feedback assessment on the sample's trace during evaluate() (default: True). Requires the traces of synalinks.enable_observability(); the assessments show up on the traces of the evaluation run in the MLflow UI.

True

Example:

import synalinks

# Basic usage - uses local MLflow storage
monitor = synalinks.callbacks.Monitor(experiment_name="my_experiment")

# With custom MLflow tracking server
monitor = synalinks.callbacks.Monitor(
    tracking_uri="http://localhost:5000",
    experiment_name="my_experiment",
    run_name="training_run_1",
    log_program_plot=True,
    log_program_model=True,
    tags={"model_type": "chain_of_thought"}
)

# Use in training
program.fit(
    x=train_data,
    y=train_labels,
    epochs=10,
    callbacks=[monitor]
)

# Track evaluation results over time: every evaluate() call (in this
# process or a later one) appends one point to the same run's charts
monitor = synalinks.callbacks.Monitor(
    experiment_name="my_experiment",
    run_name="nightly_eval",
    resume=True,
)
program.evaluate(x=test_data, y=test_labels, callbacks=[monitor])
Note

For tracing module calls along with training metrics, use synalinks.enable_observability() at the beggining of your script which configures the Monitor hook & callback:

synalinks.enable_observability(
    tracking_uri="http://localhost:5000",
    experiment_name="my_traces"
)
Source code in synalinks/src/callbacks/monitor.py
 265
 266
 267
 268
 269
 270
 271
 272
 273
 274
 275
 276
 277
 278
 279
 280
 281
 282
 283
 284
 285
 286
 287
 288
 289
 290
 291
 292
 293
 294
 295
 296
 297
 298
 299
 300
 301
 302
 303
 304
 305
 306
 307
 308
 309
 310
 311
 312
 313
 314
 315
 316
 317
 318
 319
 320
 321
 322
 323
 324
 325
 326
 327
 328
 329
 330
 331
 332
 333
 334
 335
 336
 337
 338
 339
 340
 341
 342
 343
 344
 345
 346
 347
 348
 349
 350
 351
 352
 353
 354
 355
 356
 357
 358
 359
 360
 361
 362
 363
 364
 365
 366
 367
 368
 369
 370
 371
 372
 373
 374
 375
 376
 377
 378
 379
 380
 381
 382
 383
 384
 385
 386
 387
 388
 389
 390
 391
 392
 393
 394
 395
 396
 397
 398
 399
 400
 401
 402
 403
 404
 405
 406
 407
 408
 409
 410
 411
 412
 413
 414
 415
 416
 417
 418
 419
 420
 421
 422
 423
 424
 425
 426
 427
 428
 429
 430
 431
 432
 433
 434
 435
 436
 437
 438
 439
 440
 441
 442
 443
 444
 445
 446
 447
 448
 449
 450
 451
 452
 453
 454
 455
 456
 457
 458
 459
 460
 461
 462
 463
 464
 465
 466
 467
 468
 469
 470
 471
 472
 473
 474
 475
 476
 477
 478
 479
 480
 481
 482
 483
 484
 485
 486
 487
 488
 489
 490
 491
 492
 493
 494
 495
 496
 497
 498
 499
 500
 501
 502
 503
 504
 505
 506
 507
 508
 509
 510
 511
 512
 513
 514
 515
 516
 517
 518
 519
 520
 521
 522
 523
 524
 525
 526
 527
 528
 529
 530
 531
 532
 533
 534
 535
 536
 537
 538
 539
 540
 541
 542
 543
 544
 545
 546
 547
 548
 549
 550
 551
 552
 553
 554
 555
 556
 557
 558
 559
 560
 561
 562
 563
 564
 565
 566
 567
 568
 569
 570
 571
 572
 573
 574
 575
 576
 577
 578
 579
 580
 581
 582
 583
 584
 585
 586
 587
 588
 589
 590
 591
 592
 593
 594
 595
 596
 597
 598
 599
 600
 601
 602
 603
 604
 605
 606
 607
 608
 609
 610
 611
 612
 613
 614
 615
 616
 617
 618
 619
 620
 621
 622
 623
 624
 625
 626
 627
 628
 629
 630
 631
 632
 633
 634
 635
 636
 637
 638
 639
 640
 641
 642
 643
 644
 645
 646
 647
 648
 649
 650
 651
 652
 653
 654
 655
 656
 657
 658
 659
 660
 661
 662
 663
 664
 665
 666
 667
 668
 669
 670
 671
 672
 673
 674
 675
 676
 677
 678
 679
 680
 681
 682
 683
 684
 685
 686
 687
 688
 689
 690
 691
 692
 693
 694
 695
 696
 697
 698
 699
 700
 701
 702
 703
 704
 705
 706
 707
 708
 709
 710
 711
 712
 713
 714
 715
 716
 717
 718
 719
 720
 721
 722
 723
 724
 725
 726
 727
 728
 729
 730
 731
 732
 733
 734
 735
 736
 737
 738
 739
 740
 741
 742
 743
 744
 745
 746
 747
 748
 749
 750
 751
 752
 753
 754
 755
 756
 757
 758
 759
 760
 761
 762
 763
 764
 765
 766
 767
 768
 769
 770
 771
 772
 773
 774
 775
 776
 777
 778
 779
 780
 781
 782
 783
 784
 785
 786
 787
 788
 789
 790
 791
 792
 793
 794
 795
 796
 797
 798
 799
 800
 801
 802
 803
 804
 805
 806
 807
 808
 809
 810
 811
 812
 813
 814
 815
 816
 817
 818
 819
 820
 821
 822
 823
 824
 825
 826
 827
 828
 829
 830
 831
 832
 833
 834
 835
 836
 837
 838
 839
 840
 841
 842
 843
 844
 845
 846
 847
 848
 849
 850
 851
 852
 853
 854
 855
 856
 857
 858
 859
 860
 861
 862
 863
 864
 865
 866
 867
 868
 869
 870
 871
 872
 873
 874
 875
 876
 877
 878
 879
 880
 881
 882
 883
 884
 885
 886
 887
 888
 889
 890
 891
 892
 893
 894
 895
 896
 897
 898
 899
 900
 901
 902
 903
 904
 905
 906
 907
 908
 909
 910
 911
 912
 913
 914
 915
 916
 917
 918
 919
 920
 921
 922
 923
 924
 925
 926
 927
 928
 929
 930
 931
 932
 933
 934
 935
 936
 937
 938
 939
 940
 941
 942
 943
 944
 945
 946
 947
 948
 949
 950
 951
 952
 953
 954
 955
 956
 957
 958
 959
 960
 961
 962
 963
 964
 965
 966
 967
 968
 969
 970
 971
 972
 973
 974
 975
 976
 977
 978
 979
 980
 981
 982
 983
 984
 985
 986
 987
 988
 989
 990
 991
 992
 993
 994
 995
 996
 997
 998
 999
1000
1001
1002
1003
1004
1005
1006
1007
1008
1009
1010
1011
1012
1013
1014
1015
1016
1017
1018
1019
1020
1021
1022
1023
1024
1025
1026
1027
1028
1029
1030
1031
1032
1033
1034
1035
1036
1037
1038
1039
1040
1041
1042
1043
1044
1045
1046
1047
1048
1049
1050
1051
1052
1053
1054
1055
1056
1057
1058
1059
1060
1061
1062
1063
1064
1065
1066
1067
1068
1069
1070
1071
1072
1073
1074
1075
1076
1077
1078
1079
1080
1081
1082
1083
1084
1085
1086
1087
1088
1089
1090
1091
1092
1093
1094
1095
1096
1097
1098
1099
1100
1101
1102
1103
1104
1105
1106
1107
1108
1109
1110
1111
1112
1113
1114
1115
1116
1117
1118
1119
1120
1121
1122
1123
1124
1125
1126
1127
1128
1129
1130
1131
1132
1133
1134
1135
1136
1137
1138
1139
1140
1141
1142
1143
1144
1145
1146
1147
1148
1149
1150
1151
1152
1153
1154
1155
1156
1157
1158
1159
1160
1161
1162
1163
1164
1165
1166
1167
1168
1169
1170
1171
1172
1173
1174
1175
1176
1177
1178
1179
1180
1181
1182
1183
1184
1185
1186
1187
1188
1189
1190
1191
1192
1193
1194
1195
1196
1197
1198
1199
1200
1201
1202
1203
1204
1205
1206
1207
1208
1209
1210
1211
1212
1213
1214
1215
1216
1217
1218
1219
@synalinks_export("synalinks.callbacks.Monitor")
class Monitor(Callback):
    """Monitor callback for logging training metrics to MLflow.

    This callback logs training progress and evaluation metrics to MLflow
    for experiment tracking and visualization.

    What gets logged on a `fit()` run:

    - **Metrics** per epoch (`step=epoch`): every training metric, the
      `val_*` validation metrics, the LM/EM spend of the epoch (`epoch_cost`,
      `epoch_tokens`, per phase `epoch_cost_inference|reward|optimizer`), the
      running spend of the fit (`fit_cost`, `fit_tokens`), the all-time
      in-process spend (`total_cost`, `total_tokens`) and `epoch_duration_s`;
      `fit_duration_s` at the end. With `log_batch_metrics=True`, batch
      metrics are logged too, prefixed `train_batch_` / `val_batch_` with
      their own step counters.
    - **Params**: the fit arguments (epochs, batch/minibatch size,
      validation split/freq, dataset sizes), the compile configuration
      (reward, metrics, optimizer and its config), the program (name, class,
      number of modules and trainable variables) and every language or
      embedding model reachable from the program (`lm.<name>.model`, api
      base, sampling settings).
    - **Datasets**: the train and validation sets as MLflow dataset inputs,
      one row per sample with JSON `inputs` and `expectations` columns.
    - **Model**: the program (architecture and trained variables, the JSON
      of `Program.save()`) as an MLflow pyfunc model wrapped in
      `SynalinksProgramModel`, so it can be loaded with
      `mlflow.pyfunc.load_model()` and served with `mlflow models serve`. A
      new version is logged every time `val_reward` improves (`reward` when
      there is no validation data, the end of training when neither is
      available), registered in the Model Registry under the program's name;
      the best one is tagged `synalinks.best`. Each is an MLflow LoggedModel
      linked to the run; its id is written into the saved program
      (`program._mlflow_model_id`) so later `evaluate()` runs of the same
      program link their traces and metrics to that version and nest under
      the training run.
    - **Prompts**: each trainable `Generator`-like module's prompt (the
      system turn rendered from its optimized instructions and examples, plus
      a `{{ inputs }}` user turn, the module's output schema as
      `response_format`, and its LM settings as `model_config`) is registered
      in the MLflow Prompt Registry under
      `<program>.<module>`. A new version is created at every epoch where the
      rendered prompt changed, its commit message and `val_reward` tag
      recording that epoch's validation reward; the best version gets the
      `best` alias at the end. Versions are linked to the run, to the logged
      model, and to the traces of the module's later calls.
    - **Artifacts**: the program plot at the beginning of training.

    A standalone `evaluate()` opens its own run (type `genai_evaluate`), logs
    the evaluation metrics, `eval_cost`, `eval_tokens`, `eval_duration_s`, the
    evaluated dataset, and every sample's reward as a `reward` feedback
    assessment on its trace.

    Args:
        experiment_name (str): Name of the MLflow experiment. If None, uses
            the experiment of `synalinks.enable_observability()` when it was
            called (so runs land next to their traces), else the program name.
        run_name (str): Name of the MLflow run. If None, the program's name,
            suffixed `_train` or `_test`.
        tracking_uri (str): MLflow tracking server URI. If None, uses the
            value from `synalinks.enable_observability()` or the default
            (local ./mlruns directory or MLFLOW_TRACKING_URI env var).
        log_batch_metrics (bool): Whether to log metrics at batch level
            (default: False).
        log_epoch_metrics (bool): Whether to log metrics at epoch level
            (default: True).
        log_program_plot (bool): Whether to log the program plot as an artifact
            at the beginning of training (default: True).
        log_program_model (bool): Whether to log the program as an MLflow model
            (default: True).
        tags (dict): Optional tags to add to the MLflow run.
        run_id (str): Optional. The id of an existing MLflow run to resume
            instead of starting a new one. Metrics keep being appended to it,
            the step counters continuing after the last step already logged.
        resume (bool): Whether to look up an existing run named `run_name`
            in the experiment and resume it (default: False). Creates the
            run the first time. This is how repeated `evaluate()` calls
            draw a chart over time: each evaluation adds one point to the
            metrics of the same run.
        log_assessments (bool): Whether to log each evaluated sample's reward
            as a `reward` feedback assessment on the sample's trace during
            `evaluate()` (default: True). Requires the traces of
            `synalinks.enable_observability()`; the assessments show up on
            the traces of the evaluation run in the MLflow UI.

    Example:

    ```python
    import synalinks

    # Basic usage - uses local MLflow storage
    monitor = synalinks.callbacks.Monitor(experiment_name="my_experiment")

    # With custom MLflow tracking server
    monitor = synalinks.callbacks.Monitor(
        tracking_uri="http://localhost:5000",
        experiment_name="my_experiment",
        run_name="training_run_1",
        log_program_plot=True,
        log_program_model=True,
        tags={"model_type": "chain_of_thought"}
    )

    # Use in training
    program.fit(
        x=train_data,
        y=train_labels,
        epochs=10,
        callbacks=[monitor]
    )

    # Track evaluation results over time: every evaluate() call (in this
    # process or a later one) appends one point to the same run's charts
    monitor = synalinks.callbacks.Monitor(
        experiment_name="my_experiment",
        run_name="nightly_eval",
        resume=True,
    )
    program.evaluate(x=test_data, y=test_labels, callbacks=[monitor])
    ```

    Note:
        For tracing module calls along with training metrics, use
        `synalinks.enable_observability()` at the beggining of your script
        which configures the Monitor hook & callback:

        ```python
        synalinks.enable_observability(
            tracking_uri="http://localhost:5000",
            experiment_name="my_traces"
        )
        ```
    """

    def __init__(
        self,
        experiment_name=None,
        run_name=None,
        tracking_uri=None,
        log_batch_metrics=False,
        log_epoch_metrics=True,
        log_program_plot=True,
        log_program_model=True,
        tags=None,
        run_id=None,
        resume=False,
        log_assessments=True,
    ):
        super().__init__()
        if not MLFLOW_AVAILABLE:
            raise ImportError(
                "mlflow is required for the Monitor callback. "
                "Install it with: pip install mlflow"
            )

        self.experiment_name = experiment_name
        self.run_name = run_name
        self.tracking_uri = tracking_uri or mlflow_tracking_uri()
        self.log_batch_metrics = log_batch_metrics
        self.log_epoch_metrics = log_epoch_metrics
        self.log_program_plot = log_program_plot
        self.log_program_model = log_program_model
        self.tags = tags or {}
        self.run_id = run_id
        self.resume = resume
        self.log_assessments = log_assessments
        self.logger = logging.getLogger(__name__)

        self._run = None
        self._trace_mark = None
        self._steps = {"train": 0, "val": 0, "test": 0}
        self._epoch = 0
        # Track if we're inside fit() to avoid ending run during validation
        self._in_training = False
        self._fit_counters = None
        self._epoch_counters = None
        self._fit_t0 = None
        self._epoch_t0 = None
        self._model_best = -np.inf
        self._logged_models = []
        self._active_model = None
        self._prompt_versions = {}
        self._prompt_hashes = {}
        self._prompt_history = {}
        self._prompt_snapshot = {}

    def _batch_phase(self):
        return "val" if self._in_training else "test"

    def _setup_mlflow(self):
        """Configure MLflow tracking."""
        if self.tracking_uri:
            mlflow.set_tracking_uri(self.tracking_uri)

        experiment_name = self.experiment_name
        if experiment_name is None and is_observability_enabled():
            experiment_name = mlflow_experiment_name()
        if experiment_name is None and self.program is not None:
            experiment_name = self.program.name or "synalinks_experiment"

        self._experiment_id = mlflow.set_experiment(experiment_name).experiment_id

    def _start_run(self, run_name_suffix="", tags=None):
        """Start a new MLflow run, or resume one (`run_id` / `resume`)."""
        run_name = self.run_name
        if run_name is None and self.program is not None and self.program.name:
            run_name = self.program.name
        if run_name and run_name_suffix:
            run_name = f"{run_name}_{run_name_suffix}"
        elif run_name_suffix:
            run_name = run_name_suffix

        run_id = self.run_id
        if run_id is None and self.resume:
            client = mlflow.MlflowClient()
            found = client.search_runs(
                experiment_ids=[self._experiment_id],
                filter_string=f"tags.mlflow.runName = '{run_name}'",
                order_by=["attributes.start_time DESC"],
                max_results=1,
            )
            if found:
                run_id = found[0].info.run_id

        if run_id is not None:
            self._run = mlflow.start_run(run_id=run_id)
            self.run_id = run_id
        elif tags:
            self._run = mlflow.start_run(run_name=run_name, tags=tags)
            self.run_id = self._run.info.run_id
        else:
            self._run = mlflow.start_run(run_name=run_name)
            self.run_id = self._run.info.run_id

        tags = dict(self.tags)
        if self.program is not None:
            if self.program.name:
                tags["program_name"] = self.program.name
            if self.program.description:
                tags["program_description"] = self.program.description

        if tags:
            mlflow.set_tags(tags)

        self._steps = {"train": 0, "val": 0, "test": 0}
        self._epoch = 0
        if run_id is not None:
            client = mlflow.MlflowClient()
            last_step = 0
            for key in client.get_run(run_id).data.metrics:
                history = client.get_metric_history(run_id, key)
                last_step = max(last_step, *(m.step + 1 for m in history))
            self._steps = {phase: last_step for phase in self._steps}

    def _end_run(self):
        """End the current MLflow run."""
        if self._active_model is not None:
            try:
                self._active_model.__exit__(None, None, None)
            except Exception as e:
                self.logger.debug(f"Failed to restore the active model: {e}")
            self._active_model = None
        if self._run is not None:
            mlflow.end_run()
            self._run = None

    def _activate_model(self, model_id):
        """Make `model_id` the active LoggedModel so traces link to it."""
        if model_id is None or self._active_model is not None:
            return
        try:
            self._active_model = mlflow.set_active_model(model_id=model_id)
        except Exception as e:
            self.logger.warning(f"Failed to set the active model {model_id}: {e}")

    async def _log_metrics(self, logs, step=None, prefix="", model_id=None):
        """Log metrics to MLflow asynchronously."""
        if logs is None or self._run is None:
            return

        metrics = {}
        for key, value in logs.items():
            if isinstance(value, bool):
                continue
            if isinstance(value, (int, float)):
                metrics[prefix + key] = value

        if metrics:
            kwargs = {"step": step, "run_id": self._run.info.run_id}
            if model_id is not None:
                kwargs["model_id"] = model_id
            # Explicit run_id: MLflow's active run is thread-local, not seen by the worker
            await asyncio.to_thread(mlflow.log_metrics, metrics, **kwargs)

    def _models(self):
        if self.program is None:
            return []
        return _collect_language_models(self.program) + _collect_embedding_models(
            self.program
        )

    def _snapshot_counters(self):
        """Sum the all-time and per-phase spend counters of every LM and EM."""
        models = self._models()
        counters = {}
        for suffix in _COUNTER_SUFFIXES:
            counters[suffix] = sum(getattr(m, f"cumulated_{suffix}", 0) for m in models)
            for phase in _PHASES:
                counters[f"{phase}_{suffix}"] = sum(
                    getattr(m, f"{phase}_cumulated_{suffix}", 0) for m in models
                )
        return counters

    def _spend_metrics(self, prefix, baseline, current=None):
        """Cost/token deltas since `baseline`, named `<prefix>_cost` etc."""
        if baseline is None:
            return {}
        current = current or self._snapshot_counters()
        metrics = {
            f"{prefix}_cost": current["cost"] - baseline["cost"],
            f"{prefix}_tokens": current["tokens"] - baseline["tokens"],
            f"{prefix}_calls": current["calls"] - baseline["calls"],
        }
        for phase in _PHASES:
            delta = current[f"{phase}_cost"] - baseline[f"{phase}_cost"]
            if delta:
                metrics[f"{prefix}_cost_{phase}"] = delta
        return metrics

    def _collect_params(self):
        """Fit arguments, compile configuration, program and model settings."""
        params = {}
        for key, value in (self.params or {}).items():
            if isinstance(value, (str, int, float, bool)):
                params[key] = value
        program = self.program
        if program is None:
            return params
        params["program_name"] = program.name or ""
        params["program_class"] = program.__class__.__name__
        trainable_variables = getattr(program, "trainable_variables", None)
        if trainable_variables is not None:
            params["num_trainable_variables"] = len(trainable_variables)
        if hasattr(program, "_flatten_modules"):
            params["num_modules"] = len(
                program._flatten_modules(include_self=False, recursive=True)
            )
        reward = getattr(program, "reward", None)
        if reward is not None:
            params["reward"] = getattr(reward, "name", reward.__class__.__name__)
            if hasattr(reward, "get_config"):
                params["reward_config"] = json.dumps(reward.get_config(), default=str)
        optimizer = getattr(program, "optimizer", None)
        if optimizer is not None and hasattr(optimizer, "get_config"):
            params["optimizer_config"] = json.dumps(optimizer.get_config(), default=str)
        compile_metrics = getattr(program, "_compile_metrics", None)
        metrics = getattr(compile_metrics, "metrics", None)
        if metrics:
            params["metrics"] = ",".join(m.name for m in metrics)
        for model in self._models():
            kind = "em" if "Embedding" in model.__class__.__name__ else "lm"
            key = f"{kind}.{model.name}"
            params[f"{key}.model"] = getattr(model, "model", "")
            api_base = getattr(model, "api_base", None)
            if api_base:
                params[f"{key}.api_base"] = api_base
            for name, value in (getattr(model, "default_kwargs", None) or {}).items():
                if isinstance(value, (str, int, float, bool)):
                    params[f"{key}.{name}"] = value
        for key, value in list(params.items()):
            if isinstance(value, str) and len(value) > _MAX_PARAM_LENGTH:
                params[key] = value[:_MAX_PARAM_LENGTH]
        return params

    async def _log_params(self):
        """Log the run parameters to MLflow asynchronously.

        Params are immutable in MLflow: resuming a run with a changed value
        raises, so on failure fall back to logging key by key and skip the
        conflicting ones.
        """
        if self._run is None:
            return
        run_id = self._run.info.run_id
        params = self._collect_params()
        if not params:
            return
        try:
            await asyncio.to_thread(mlflow.log_params, params, run_id=run_id)
        except Exception as e:
            self.logger.debug(f"Batch log_params failed ({e}), logging key by key")
            client = mlflow.MlflowClient()
            for key, value in params.items():
                try:
                    await asyncio.to_thread(client.log_param, run_id, key, value)
                except Exception as err:
                    self.logger.warning(f"Failed to log param {key}: {err}")

    async def _log_dataset(self, x, y, context):
        """Log `x`/`y` as an MLflow dataset input of the run.

        Rows hold JSON strings (not dicts) so MLflow's pandas digest covers
        the data. Only the schema, digest and size are stored on the run.
        """
        if self._run is None or x is None:
            return
        try:
            import pandas as pd
            from mlflow.entities import DatasetInput
            from mlflow.entities import InputTag

            records = {"inputs": [json.dumps(item.get_json()) for item in x]}
            if y is not None:
                records["expectations"] = [json.dumps(item.get_json()) for item in y]
            df = pd.DataFrame(records)
            name = f"{self.program.name or 'program'}_{context}"
            dataset = mlflow.data.from_pandas(
                df, name=name, targets="expectations" if y is not None else None
            )
            client = mlflow.MlflowClient()
            await asyncio.to_thread(
                client.log_inputs,
                self._run.info.run_id,
                datasets=[
                    DatasetInput(
                        dataset._to_mlflow_entity(),
                        tags=[InputTag("mlflow.data.context", context)],
                    )
                ],
            )
        except Exception as e:
            self.logger.warning(f"Failed to log {context} dataset: {e}")

    async def _log_artifact(self, local_path, artifact_path):
        await asyncio.to_thread(
            mlflow.log_artifact,
            local_path,
            artifact_path=artifact_path,
            run_id=self._run.info.run_id,
        )

    async def _log_program_plot_artifact(self):
        """Log the program plot as an MLflow artifact asynchronously."""
        if self._run is None:
            self.logger.warning("No MLflow run active, skipping plot logging")
            return

        if self.program is None:
            self.logger.warning("No program set, skipping plot logging")
            return

        if not self.program.built:
            self.logger.warning("Program not built, skipping plot logging")
            return

        try:
            from synalinks.src.utils.program_visualization import check_graphviz
            from synalinks.src.utils.program_visualization import check_pydot
            from synalinks.src.utils.program_visualization import plot_program

            if not check_pydot() or not check_graphviz():
                self.logger.warning(
                    "pydot or graphviz not available, skipping program plot"
                )
                return

            with tempfile.TemporaryDirectory() as tmpdir:
                plot_filename = f"{self.program.name or 'program'}.png"
                plot_path = os.path.join(tmpdir, plot_filename)

                await asyncio.to_thread(
                    plot_program,
                    self.program,
                    to_file=plot_filename,
                    to_folder=tmpdir,
                    show_schemas=True,
                    show_module_names=True,
                    show_trainable=True,
                    dpi=96,
                )

                if os.path.exists(plot_path):
                    await self._log_artifact(plot_path, "program_plots")
                    self.logger.info(f"Logged program plot: {plot_filename}")
                else:
                    self.logger.warning(f"Plot file not created: {plot_path}")

        except Exception as e:
            self.logger.warning(f"Failed to log program plot: {e}")

    def _prompt_modules(self):
        if self.program is None or not hasattr(self.program, "_flatten_modules"):
            return []
        return [
            m
            for m in self.program._flatten_modules(include_self=False, recursive=True)
            if hasattr(m, "_render_system_message") and hasattr(m, "state")
        ]

    def _snapshot_prompts(self):
        """Render every trainable prompt: `{name: (hash, messages, module)}`."""
        snapshot = {}
        for module in self._prompt_modules():
            try:
                messages = _render_prompt_messages(module)
            except Exception as e:
                self.logger.warning(f"Failed to render the prompt of {module.name}: {e}")
                continue
            if messages is None:
                continue
            digest = hashlib.sha256(
                json.dumps(messages, sort_keys=True).encode()
            ).hexdigest()
            name = _prompt_name(self.program.name, module.name)
            snapshot[name] = (digest, messages, module)
        return snapshot

    async def _register_prompts(self, epoch, logs):
        """Register a new version of every prompt that changed this epoch."""
        if self._run is None or not self._prompt_snapshot:
            return
        value = self._monitored_value(logs)
        run_id = self._run.info.run_id
        model_id = getattr(self.program, "_mlflow_model_id", None)
        optimizer = (self.params or {}).get("optimizer")
        client = mlflow.MlflowClient()
        for name, (digest, messages, module) in self._prompt_snapshot.items():
            if self._prompt_hashes.get(name) == digest:
                continue
            try:
                prompt_version = None
                if name not in self._prompt_hashes:
                    existing = await asyncio.to_thread(
                        client.load_prompt, name, allow_missing=True
                    )
                    if existing is not None and existing.template == messages:
                        prompt_version = existing
                if prompt_version is None:
                    tags = {
                        "program": self.program.name or "",
                        "module": module.name or "",
                        "epoch": str(epoch),
                        "run_id": run_id,
                    }
                    if value is not None:
                        tags["val_reward"] = str(value)
                    if optimizer:
                        tags["optimizer"] = str(optimizer)
                    commit_message = f"epoch {epoch}"
                    if value is not None:
                        commit_message += f": val_reward={value:.4f}"
                    prompt_version = await asyncio.to_thread(
                        mlflow.genai.register_prompt,
                        name=name,
                        template=messages,
                        commit_message=commit_message,
                        tags=tags,
                        response_format=getattr(module, "schema", None),
                        model_config=_prompt_model_config(module),
                    )
                await asyncio.to_thread(
                    client.link_prompt_version_to_run, run_id, prompt_version
                )
                if model_id:
                    await asyncio.to_thread(
                        client.link_prompt_version_to_model,
                        name,
                        str(prompt_version.version),
                        model_id,
                    )
                self._prompt_hashes[name] = digest
                self._prompt_versions[name] = prompt_version
                self._prompt_history.setdefault(name, []).append(
                    (prompt_version.version, value)
                )
                module._mlflow_prompt_version = prompt_version
            except Exception as e:
                self.logger.warning(f"Failed to register prompt {name}: {e}")

    def _set_prompt_aliases(self):
        """Point the `best` alias of each prompt at its best-scoring version."""
        for name, history in self._prompt_history.items():
            version = history[-1][0]
            scored = [h for h in history if h[1] is not None]
            if scored:
                version = max(scored, key=lambda h: h[1])[0]
            try:
                mlflow.genai.set_prompt_alias(name, "best", version)
            except Exception as e:
                self.logger.debug(f"Failed to set the best alias of {name}: {e}")

    def _monitored_value(self, logs):
        """The epoch's `val_reward`, else `reward`, else `None`."""
        logs = logs or {}
        value = logs.get("val_reward")
        if value is None:
            value = logs.get("reward")
        return value

    def _maybe_log_model(self, epoch, logs):
        """Log a new model version when the reward improved this epoch."""
        if not self.log_program_model:
            return
        value = self._monitored_value(logs)
        if value is None or value <= self._model_best:
            return
        self._model_best = value
        self._log_model(step=epoch, value=value)

    def _input_example(self):
        inputs = getattr(self.program, "_fit_inputs", None) or {}
        x = inputs.get("x")
        if x is None or len(x) == 0:
            return None
        return x[0].get_json()

    def _log_model(self, step, value=None):
        """Log the program as a pyfunc model and a LoggedModel version.

        Two-phase so the saved `program.json` carries its own model id: the
        LoggedModel is created first, its id written on the program, then the
        model files are logged against it. Runs on the callback thread since
        MLflow resolves the active run thread-locally.
        """
        if self._run is None or self.program is None:
            return None

        run_id = self._run.info.run_id
        name = self.program.name or "program"
        tags = {"synalinks.program": name, "synalinks.epoch": str(step)}
        if value is not None:
            tags["synalinks.reward"] = str(value)
        try:
            logged_model = mlflow.initialize_logged_model(
                name=name, source_run_id=run_id, model_type=_MODEL_TYPE, tags=tags
            )
            model_id = logged_model.model_id
            self.program._mlflow_model_id = model_id
            self.program._mlflow_run_id = run_id
            self.program._mlflow_experiment_id = getattr(self, "_experiment_id", None)
            self.program._mlflow_prompts = {
                module.name: {
                    "name": pv.name,
                    "version": pv.version,
                    "uri": getattr(pv, "uri", None),
                }
                for module in self._prompt_modules()
                for pv in [getattr(module, "_mlflow_prompt_version", None)]
                if pv is not None
            }
            input_example = self._input_example()
            signature = build_signature(
                getattr(self.program, "input_schema", None),
                getattr(self.program, "output_schema", None),
                input_example=input_example,
            )
            with tempfile.TemporaryDirectory() as tmpdir:
                path = os.path.join(tmpdir, "program.json")
                self.program.save(path)
                info = mlflow.models.Model.log(
                    artifact_path=None,
                    flavor=mlflow.pyfunc,
                    name=name,
                    run_id=run_id,
                    model_id=model_id,
                    model_type=_MODEL_TYPE,
                    step=step,
                    tags=tags,
                    prompts=[pv.uri for pv in self._prompt_versions.values()] or None,
                    registered_model_name=name,
                    python_model=SynalinksProgramModel(),
                    artifacts={"program": path},
                    signature=signature,
                    input_example=input_example,
                    pip_requirements=[f"synalinks=={__version__}"],
                )
            mlflow.finalize_logged_model(model_id, "READY")
            self._logged_models.append((step, model_id, value))
            self._activate_model(model_id)
            self.logger.info(f"Logged program model {name} (model_id={model_id})")
            return info
        except Exception as e:
            self.logger.warning(f"Failed to log program model: {e}")
            return None

    def _tag_best_model(self):
        if not self._logged_models:
            return
        best = self._logged_models[-1]
        scored = [m for m in self._logged_models if m[2] is not None]
        if scored:
            best = max(scored, key=lambda m: m[2])
        try:
            mlflow.MlflowClient().set_logged_model_tags(
                best[1], {"synalinks.best": "true"}
            )
        except Exception as e:
            self.logger.debug(f"Failed to tag the best model: {e}")

    def on_train_begin(self, logs=None):
        """Called at the beginning of training."""
        self._in_training = True
        self._setup_mlflow()
        self._start_run(run_name_suffix="train")
        self.logger.debug("MLflow run started for training")
        self._fit_t0 = time.perf_counter()
        self._fit_counters = self._snapshot_counters()
        self._logged_models = []
        self._model_best = -np.inf
        self._prompt_versions = {}
        self._prompt_hashes = {}
        self._prompt_history = {}
        self._prompt_snapshot = {}

        run_maybe_nested(self._log_params())

        inputs = getattr(self.program, "_fit_inputs", None) or {}
        run_maybe_nested(self._log_dataset(inputs.get("x"), inputs.get("y"), "training"))
        run_maybe_nested(
            self._log_dataset(inputs.get("val_x"), inputs.get("val_y"), "validation")
        )

        if self.log_program_plot:
            run_maybe_nested(self._log_program_plot_artifact())

    def on_train_end(self, logs=None):
        """Called at the end of training."""
        if self._fit_t0 is not None:
            run_maybe_nested(
                self._log_metrics(
                    {"fit_duration_s": time.perf_counter() - self._fit_t0},
                    step=self._epoch,
                )
            )

        if self.log_program_model and not self._logged_models:
            self._log_model(step=self._epoch, value=self._monitored_value(logs))
        self._tag_best_model()
        self._set_prompt_aliases()

        self._end_run()
        self._in_training = False
        self.logger.debug("MLflow run ended for training")

    def on_epoch_begin(self, epoch, logs=None):
        """Called at the start of an epoch."""
        self._epoch = epoch
        self._epoch_t0 = time.perf_counter()
        self._epoch_counters = self._snapshot_counters()

    def on_epoch_end(self, epoch, logs=None):
        """Called at the end of an epoch."""
        self._epoch = epoch
        if not self.log_epoch_metrics:
            return

        metrics = dict(logs or {})
        current = self._snapshot_counters()
        metrics.update(self._spend_metrics("epoch", self._epoch_counters, current))
        metrics.update(self._spend_metrics("fit", self._fit_counters, current))
        metrics["total_cost"] = current["cost"]
        metrics["total_tokens"] = current["tokens"]
        if self._epoch_t0 is not None:
            metrics["epoch_duration_s"] = time.perf_counter() - self._epoch_t0
        run_maybe_nested(self._log_metrics(metrics, step=epoch))
        self.logger.debug(f"Logged metrics for epoch {epoch}")
        # Snapshot taken when validation began; without validation the
        # state at epoch end is the one to register.
        if not self._prompt_snapshot:
            self._prompt_snapshot = self._snapshot_prompts()
        run_maybe_nested(self._register_prompts(epoch, logs))
        self._prompt_snapshot = {}
        self._maybe_log_model(epoch, logs)

    def on_train_batch_begin(self, batch, logs=None):
        """Called at the beginning of a training batch."""
        pass

    def on_train_batch_end(self, batch, logs=None):
        """Called at the end of a training batch."""
        if not self.log_batch_metrics:
            return

        step = self._steps["train"]
        self._steps["train"] += 1
        run_maybe_nested(self._log_metrics(logs, step=step, prefix="train_batch_"))

    def on_test_begin(self, logs=None):
        """Called at the beginning of evaluation or validation."""
        if self._in_training:
            # The state being validated is the one whose reward we record.
            self._prompt_snapshot = self._snapshot_prompts()
        # Only start a new run if we're not already in a training run
        if self._run is None and not self._in_training:
            self._setup_mlflow()
            tags = None
            parent_run_id = getattr(self.program, "_mlflow_run_id", None)
            same_experiment = (
                getattr(self.program, "_mlflow_experiment_id", None)
                == self._experiment_id
            )
            if parent_run_id and same_experiment and not (self.run_id or self.resume):
                tags = {"mlflow.parentRunId": parent_run_id}
            self._start_run(run_name_suffix="test", tags=tags)
            self._activate_model(getattr(self.program, "_mlflow_model_id", None))
            # Same run type as `mlflow.genai.evaluate()` runs
            mlflow.set_tag("mlflow.runType", "genai_evaluate")
            self.logger.debug("MLflow run started for testing")
            self._eval_t0 = time.perf_counter()
            self._eval_counters = self._snapshot_counters()
            run_maybe_nested(self._log_params())
            inputs = getattr(self.program, "_eval_inputs", None) or {}
            run_maybe_nested(
                self._log_dataset(inputs.get("x"), inputs.get("y"), "evaluation")
            )

    def on_test_end(self, logs=None):
        """Called at the end of evaluation or validation.

        Inside `fit()` the trainer merges the validation metrics, prefixed
        `val_`, into the epoch logs, so nothing is logged here: logging the
        unprefixed values would overwrite the training metrics.
        """
        if self._in_training or self._run is None:
            return
        metrics = dict(logs or {})
        metrics.update(self._spend_metrics("eval", self._eval_counters))
        if self._eval_t0 is not None:
            metrics["eval_duration_s"] = time.perf_counter() - self._eval_t0
        run_maybe_nested(
            self._log_metrics(
                metrics,
                step=self._steps["test"],
                model_id=getattr(self.program, "_mlflow_model_id", None),
            )
        )
        self._steps["test"] += 1
        self._end_run()
        self.logger.debug("MLflow run ended for testing")

    def on_test_batch_begin(self, batch, logs=None):
        """Called at the beginning of a test batch."""
        self._trace_mark = monitor_hook.root_trace_mark()

    def on_test_batch_end(self, batch, logs=None):
        """Called at the end of a test batch."""
        if self.log_assessments:
            run_maybe_nested(self._log_batch_assessments())

        if not self.log_batch_metrics:
            return

        phase = self._batch_phase()
        step = self._steps[phase]
        self._steps[phase] += 1
        run_maybe_nested(self._log_metrics(logs, step=step, prefix=f"{phase}_batch_"))

    async def _log_batch_assessments(self):
        """Log the per-sample rewards of the batch just evaluated as `reward`
        feedback assessments on the samples' traces.

        Sample i is matched with the i-th root trace started since
        `on_test_batch_begin`; when the counts differ (tracing disabled, or
        the batch's predictions came from the auto-build pass that ran before
        the run started) nothing is logged.
        """
        if self._run is None or self.program is None:
            return
        rewards = getattr(self.program, "_per_sample_rewards", None)
        trace_ids = monitor_hook.root_trace_ids_since(self._trace_mark)
        if not rewards or len(rewards) != len(trace_ids):
            self.logger.debug(
                "Skipping assessments: %s rewards for %s traces",
                None if rewards is None else len(rewards),
                len(trace_ids),
            )
            return

        reward_fn = getattr(self.program, "_compile_reward", None)
        reward_fn = getattr(reward_fn, "_user_reward", reward_fn)
        source_id = getattr(reward_fn, "name", None) or "reward"
        source = mlflow.entities.AssessmentSource(
            source_type=mlflow.entities.AssessmentSourceType.CODE,
            source_id=source_id,
        )
        run_id = self._run.info.run_id
        targets = getattr(self.program, "_per_sample_targets", None)
        if not targets or len(targets) != len(trace_ids):
            targets = [None] * len(trace_ids)
        expectation_source = mlflow.entities.AssessmentSource(
            source_type=mlflow.entities.AssessmentSourceType.HUMAN,
            source_id="dataset",
        )

        def log_one(trace_id, value, target):
            try:
                mlflow.log_feedback(
                    trace_id=trace_id,
                    name="reward",
                    value=float(value),
                    source=source,
                    metadata={"mlflow.assessment.sourceRunId": run_id},
                )
                if target is not None:
                    mlflow.log_expectation(
                        trace_id=trace_id,
                        name="expected_output",
                        value=target,
                        source=expectation_source,
                        metadata={"mlflow.assessment.sourceRunId": run_id},
                    )
            except Exception as e:
                self.logger.warning(f"Failed to log assessment on {trace_id}: {e}")

        await asyncio.gather(
            *(
                asyncio.to_thread(log_one, trace_id, value, target)
                for trace_id, value, target in zip(trace_ids, rewards, targets)
            )
        )

    def on_predict_begin(self, logs=None):
        """Called at the beginning of prediction."""
        pass

    def on_predict_end(self, logs=None):
        """Called at the end of prediction."""
        pass

    def on_predict_batch_begin(self, batch, logs=None):
        """Called at the beginning of a prediction batch."""
        pass

    def on_predict_batch_end(self, batch, logs=None):
        """Called at the end of a prediction batch."""
        pass

    def __del__(self):
        """End our MLflow run if it was left open and is still the active one.

        Guarded by a run-id check: ``mlflow.end_run()`` always ends whatever run
        is *globally* active, so a finalizer firing at GC time must not end an
        unrelated run (this also keeps a leaked finalizer from polluting other
        code's, or another test's, active run).
        """
        run = getattr(self, "_run", None)
        if run is None:
            return
        try:
            active = mlflow.active_run()
            if active is not None and active.info.run_id == run.info.run_id:
                mlflow.end_run()
        except Exception:
            pass

__del__()

End our MLflow run if it was left open and is still the active one.

Guarded by a run-id check: mlflow.end_run() always ends whatever run is globally active, so a finalizer firing at GC time must not end an unrelated run (this also keeps a leaked finalizer from polluting other code's, or another test's, active run).

Source code in synalinks/src/callbacks/monitor.py
def __del__(self):
    """End our MLflow run if it was left open and is still the active one.

    Guarded by a run-id check: ``mlflow.end_run()`` always ends whatever run
    is *globally* active, so a finalizer firing at GC time must not end an
    unrelated run (this also keeps a leaked finalizer from polluting other
    code's, or another test's, active run).
    """
    run = getattr(self, "_run", None)
    if run is None:
        return
    try:
        active = mlflow.active_run()
        if active is not None and active.info.run_id == run.info.run_id:
            mlflow.end_run()
    except Exception:
        pass

on_epoch_begin(epoch, logs=None)

Called at the start of an epoch.

Source code in synalinks/src/callbacks/monitor.py
def on_epoch_begin(self, epoch, logs=None):
    """Called at the start of an epoch."""
    self._epoch = epoch
    self._epoch_t0 = time.perf_counter()
    self._epoch_counters = self._snapshot_counters()

on_epoch_end(epoch, logs=None)

Called at the end of an epoch.

Source code in synalinks/src/callbacks/monitor.py
def on_epoch_end(self, epoch, logs=None):
    """Called at the end of an epoch."""
    self._epoch = epoch
    if not self.log_epoch_metrics:
        return

    metrics = dict(logs or {})
    current = self._snapshot_counters()
    metrics.update(self._spend_metrics("epoch", self._epoch_counters, current))
    metrics.update(self._spend_metrics("fit", self._fit_counters, current))
    metrics["total_cost"] = current["cost"]
    metrics["total_tokens"] = current["tokens"]
    if self._epoch_t0 is not None:
        metrics["epoch_duration_s"] = time.perf_counter() - self._epoch_t0
    run_maybe_nested(self._log_metrics(metrics, step=epoch))
    self.logger.debug(f"Logged metrics for epoch {epoch}")
    # Snapshot taken when validation began; without validation the
    # state at epoch end is the one to register.
    if not self._prompt_snapshot:
        self._prompt_snapshot = self._snapshot_prompts()
    run_maybe_nested(self._register_prompts(epoch, logs))
    self._prompt_snapshot = {}
    self._maybe_log_model(epoch, logs)

on_predict_batch_begin(batch, logs=None)

Called at the beginning of a prediction batch.

Source code in synalinks/src/callbacks/monitor.py
def on_predict_batch_begin(self, batch, logs=None):
    """Called at the beginning of a prediction batch."""
    pass

on_predict_batch_end(batch, logs=None)

Called at the end of a prediction batch.

Source code in synalinks/src/callbacks/monitor.py
def on_predict_batch_end(self, batch, logs=None):
    """Called at the end of a prediction batch."""
    pass

on_predict_begin(logs=None)

Called at the beginning of prediction.

Source code in synalinks/src/callbacks/monitor.py
def on_predict_begin(self, logs=None):
    """Called at the beginning of prediction."""
    pass

on_predict_end(logs=None)

Called at the end of prediction.

Source code in synalinks/src/callbacks/monitor.py
def on_predict_end(self, logs=None):
    """Called at the end of prediction."""
    pass

on_test_batch_begin(batch, logs=None)

Called at the beginning of a test batch.

Source code in synalinks/src/callbacks/monitor.py
def on_test_batch_begin(self, batch, logs=None):
    """Called at the beginning of a test batch."""
    self._trace_mark = monitor_hook.root_trace_mark()

on_test_batch_end(batch, logs=None)

Called at the end of a test batch.

Source code in synalinks/src/callbacks/monitor.py
def on_test_batch_end(self, batch, logs=None):
    """Called at the end of a test batch."""
    if self.log_assessments:
        run_maybe_nested(self._log_batch_assessments())

    if not self.log_batch_metrics:
        return

    phase = self._batch_phase()
    step = self._steps[phase]
    self._steps[phase] += 1
    run_maybe_nested(self._log_metrics(logs, step=step, prefix=f"{phase}_batch_"))

on_test_begin(logs=None)

Called at the beginning of evaluation or validation.

Source code in synalinks/src/callbacks/monitor.py
def on_test_begin(self, logs=None):
    """Called at the beginning of evaluation or validation."""
    if self._in_training:
        # The state being validated is the one whose reward we record.
        self._prompt_snapshot = self._snapshot_prompts()
    # Only start a new run if we're not already in a training run
    if self._run is None and not self._in_training:
        self._setup_mlflow()
        tags = None
        parent_run_id = getattr(self.program, "_mlflow_run_id", None)
        same_experiment = (
            getattr(self.program, "_mlflow_experiment_id", None)
            == self._experiment_id
        )
        if parent_run_id and same_experiment and not (self.run_id or self.resume):
            tags = {"mlflow.parentRunId": parent_run_id}
        self._start_run(run_name_suffix="test", tags=tags)
        self._activate_model(getattr(self.program, "_mlflow_model_id", None))
        # Same run type as `mlflow.genai.evaluate()` runs
        mlflow.set_tag("mlflow.runType", "genai_evaluate")
        self.logger.debug("MLflow run started for testing")
        self._eval_t0 = time.perf_counter()
        self._eval_counters = self._snapshot_counters()
        run_maybe_nested(self._log_params())
        inputs = getattr(self.program, "_eval_inputs", None) or {}
        run_maybe_nested(
            self._log_dataset(inputs.get("x"), inputs.get("y"), "evaluation")
        )

on_test_end(logs=None)

Called at the end of evaluation or validation.

Inside fit() the trainer merges the validation metrics, prefixed val_, into the epoch logs, so nothing is logged here: logging the unprefixed values would overwrite the training metrics.

Source code in synalinks/src/callbacks/monitor.py
def on_test_end(self, logs=None):
    """Called at the end of evaluation or validation.

    Inside `fit()` the trainer merges the validation metrics, prefixed
    `val_`, into the epoch logs, so nothing is logged here: logging the
    unprefixed values would overwrite the training metrics.
    """
    if self._in_training or self._run is None:
        return
    metrics = dict(logs or {})
    metrics.update(self._spend_metrics("eval", self._eval_counters))
    if self._eval_t0 is not None:
        metrics["eval_duration_s"] = time.perf_counter() - self._eval_t0
    run_maybe_nested(
        self._log_metrics(
            metrics,
            step=self._steps["test"],
            model_id=getattr(self.program, "_mlflow_model_id", None),
        )
    )
    self._steps["test"] += 1
    self._end_run()
    self.logger.debug("MLflow run ended for testing")

on_train_batch_begin(batch, logs=None)

Called at the beginning of a training batch.

Source code in synalinks/src/callbacks/monitor.py
def on_train_batch_begin(self, batch, logs=None):
    """Called at the beginning of a training batch."""
    pass

on_train_batch_end(batch, logs=None)

Called at the end of a training batch.

Source code in synalinks/src/callbacks/monitor.py
def on_train_batch_end(self, batch, logs=None):
    """Called at the end of a training batch."""
    if not self.log_batch_metrics:
        return

    step = self._steps["train"]
    self._steps["train"] += 1
    run_maybe_nested(self._log_metrics(logs, step=step, prefix="train_batch_"))

on_train_begin(logs=None)

Called at the beginning of training.

Source code in synalinks/src/callbacks/monitor.py
def on_train_begin(self, logs=None):
    """Called at the beginning of training."""
    self._in_training = True
    self._setup_mlflow()
    self._start_run(run_name_suffix="train")
    self.logger.debug("MLflow run started for training")
    self._fit_t0 = time.perf_counter()
    self._fit_counters = self._snapshot_counters()
    self._logged_models = []
    self._model_best = -np.inf
    self._prompt_versions = {}
    self._prompt_hashes = {}
    self._prompt_history = {}
    self._prompt_snapshot = {}

    run_maybe_nested(self._log_params())

    inputs = getattr(self.program, "_fit_inputs", None) or {}
    run_maybe_nested(self._log_dataset(inputs.get("x"), inputs.get("y"), "training"))
    run_maybe_nested(
        self._log_dataset(inputs.get("val_x"), inputs.get("val_y"), "validation")
    )

    if self.log_program_plot:
        run_maybe_nested(self._log_program_plot_artifact())

on_train_end(logs=None)

Called at the end of training.

Source code in synalinks/src/callbacks/monitor.py
def on_train_end(self, logs=None):
    """Called at the end of training."""
    if self._fit_t0 is not None:
        run_maybe_nested(
            self._log_metrics(
                {"fit_duration_s": time.perf_counter() - self._fit_t0},
                step=self._epoch,
            )
        )

    if self.log_program_model and not self._logged_models:
        self._log_model(step=self._epoch, value=self._monitored_value(logs))
    self._tag_best_model()
    self._set_prompt_aliases()

    self._end_run()
    self._in_training = False
    self.logger.debug("MLflow run ended for training")

SynalinksProgramModel

Bases: PythonModel if MLFLOW_AVAILABLE else object

MLflow pyfunc wrapper around a saved Synalinks program.

Logged by synalinks.callbacks.Monitor at the end of training, loadable with mlflow.pyfunc.load_model(...) and servable with mlflow models serve. The program artifact is the program.json written by Program.save(); inputs are the program's input JSON (one dict, a list of dicts or a DataFrame with one column per field) and outputs are the program's output JSON dicts (None for a failed sample).

Custom DataModel, Module or Program subclasses are resolved by Program.load() from their import path, so the code declaring them only needs to be importable where the model is loaded.

Source code in synalinks/src/callbacks/monitor.py
@synalinks_export("synalinks.callbacks.SynalinksProgramModel")
class SynalinksProgramModel(mlflow.pyfunc.PythonModel if MLFLOW_AVAILABLE else object):
    """MLflow pyfunc wrapper around a saved Synalinks program.

    Logged by `synalinks.callbacks.Monitor` at the end of training, loadable
    with `mlflow.pyfunc.load_model(...)` and servable with `mlflow models
    serve`. The `program` artifact is the `program.json` written by
    `Program.save()`; inputs are the program's input JSON (one dict, a list
    of dicts or a DataFrame with one column per field) and outputs are the
    program's output JSON dicts (`None` for a failed sample).

    Custom `DataModel`, `Module` or `Program` subclasses are resolved by
    `Program.load()` from their import path, so the code declaring them only
    needs to be importable where the model is loaded.
    """

    # The input and output are validated against the model signature that
    # `Monitor` logs, derived from the program's own JSON schemas: a program's
    # inputs are arbitrary JSON, which no static type hint on `predict` can
    # describe (and a hint would reject the dict and DataFrame inputs). This
    # tells MLflow not to expect type-hint based validation.
    _skip_type_hint_validation = True

    def __init__(self):
        self.program = None
        self._loop = None

    def load_context(self, context):
        from synalinks.src.programs import Program

        self.program = Program.load(context.artifacts["program"])
        self._loop = _BackgroundLoop()

    def predict(self, context, model_input, params=None):
        records = _to_records(model_input)
        input_schema = getattr(self.program, "input_schema", None)
        x = np.array(
            [JsonDataModel(json=record, schema=input_schema) for record in records],
            dtype="object",
        )
        outputs = self._loop.run(self.program.predict(x, verbose=0))
        return [o.get_json() if o is not None else None for o in outputs]

build_signature(input_schema, output_schema, input_example=None, output_example=None)

Build the MLflow model signature of a program.

Derived from the program's input and output JSON schemas; when that fails, inferred from the examples; None when neither is possible.

Source code in synalinks/src/callbacks/monitor.py
def build_signature(input_schema, output_schema, input_example=None, output_example=None):
    """Build the MLflow model signature of a program.

    Derived from the program's input and output JSON schemas; when that
    fails, inferred from the examples; `None` when neither is possible.
    """
    from mlflow.models import ModelSignature
    from mlflow.models import infer_signature

    try:
        inputs = _json_schema_to_mlflow_schema(input_schema or {})
        outputs = _json_schema_to_mlflow_schema(output_schema or {})
        if inputs is not None:
            return ModelSignature(inputs=inputs, outputs=outputs)
    except Exception:
        pass
    try:
        if input_example is not None:
            return infer_signature(input_example, output_example)
    except Exception:
        pass
    return None