nvflare.app_opt.hf.api module

hf_is_running() → bool[source]
patch(trainer, restore_state=True, load_state_dict_strict=True, params_scope='auto', server_key_prefix=None, local_epochs=None, local_steps=None, stream_metrics=False)[source]

Patches a HuggingFace Trainer for usage with NVFlare.

Patching wraps trainer.train() and trainer.evaluate() so the usual HuggingFace training loop can execute NVFlare train, evaluate, and submit_model tasks. In distributed runs, rank 0 owns the NVFlare Client API receive/send calls and broadcasts task state to the other ranks.

Parameters:
  • trainer – HuggingFace transformers.Trainer or subclass, such as TRL SFTTrainer, to patch.

  • restore_state – Whether to resume HuggingFace optimizer and learning-rate scheduler state from the previous in-process checkpoint between FL train rounds. Defaults to True.

  • load_state_dict_strict – Exposes the strict argument of torch.nn.Module.load_state_dict() when loading received model weights. Defaults to True.

  • params_scope – Which model parameters participate in FL. Use "auto" to infer full-model versus PEFT adapter parameters, "model" for full model weights, or "adapter" for PEFT adapter weights.

  • server_key_prefix – Optional key prefix expected by the FL server. Incoming server parameters have this prefix removed before loading, and outgoing client parameters have it added before sending.

  • local_epochs – Number of local epochs per FL train round. Mutually exclusive with local_steps.

  • local_steps – Number of local optimizer steps per FL train round. Mutually exclusive with local_epochs.

  • stream_metrics – Whether to stream HuggingFace logging metrics through the NVFlare Client API metrics writer. Defaults to False.

Returns:

The patched trainer.

Raises:
  • TypeError – If trainer is not a HuggingFace Trainer.

  • ValueError – If the Trainer config is unsupported, including DeepSpeed, FSDP, load_best_model_at_end=True, save_only_model=True with restore_state=True, prebuilt optimizer/scheduler instances with restore_state=False, or both local_epochs and local_steps.

  • RuntimeError – If distributed execution is misconfigured, if restore_state=True is used with an explicit launch_once=False Client API configuration, or if another Trainer is already patched in the same process.