accelerate API reference
81 public APIs from accelerate (huggingface/accelerate) — 27 classes, 26 functions, 28 methods. Signatures extracted by static analysis of the actual source.
Repository: huggingface/accelerate
| Kind | Count |
|---|---|
| Classes | 27 |
| Functions | 26 |
| Methods | 28 |
API list
method
src.accelerate.accelerator.Accelerator.context_parallel_rank() -> intContext parallelism is not supported yet.
method
src.accelerate.accelerator.Accelerator.is_main_process()True for one process only.
method
src.accelerate.accelerator.Accelerator.lomo_backward(loss:torch.Tensor, learning_rate:float) -> NoneRuns backward pass on LOMO optimizers.
method
src.accelerate.accelerator.Accelerator.pipeline_parallel_rank() -> intPipeline parallelism is not supported yet.
method
src.accelerate.accelerator.Accelerator.tensor_parallel_rank() -> intReturns the local rank for tensor parallelism.
class
src.accelerate.commands.menu.input.KeyHandlerMetaclass that adds the key handlers to the class
func
src.accelerate.commands.menu.keymap.get_raw_chars()Gets raw characters from inputs
func
src.accelerate.hooks.add_hook_to_module(module:nn.Module, hook:ModelHook, append:bool=False)Adds a hook to a given module.
class
src.accelerate.local_sgd.LocalSGDA helper class to support local SGD on top of Accelerator.
class
src.accelerate.logging.MultiProcessAdapterAn adapter to assist with logging in multiprocess.
class
src.accelerate.optimizer.AcceleratedOptimizerInternal wrapper around a torch optimizer.
method
src.accelerate.optimizer.AcceleratedOptimizer.eval()Sets the optimizer to "eval" mode.
class
src.accelerate.tracking.AimTrackerA `Tracker` class that supports `aim`.
method
src.accelerate.tracking.AimTracker.finish()Closes `aim` writer
method
src.accelerate.tracking.AimTracker.log(values:dict, step:Optional[int]=None, **kwargs)Logs `values` to the current run.
class
src.accelerate.tracking.ClearMLTrackerA `Tracker` class that supports `clearml`.
method
src.accelerate.tracking.ClearMLTracker.finish()Close the ClearML task.
method
src.accelerate.tracking.ClearMLTracker.log_images(values:dict, step:Optional[int]=None, **kwargs)Logs `images` to the current run.
class
src.accelerate.tracking.CometMLTrackerA `Tracker` class that supports `comet_ml`.
method
src.accelerate.tracking.CometMLTracker.finish()Flush `comet-ml` writer
method
src.accelerate.tracking.CometMLTracker.log(values:dict, step:Optional[int]=None, **kwargs)Logs `values` to the current run.
class
src.accelerate.tracking.DVCLiveTrackerA `Tracker` class that supports `dvclive`.
method
src.accelerate.tracking.DVCLiveTracker.finish()Closes `dvclive.Live()`.
method
src.accelerate.tracking.DVCLiveTracker.log(values:dict, step:Optional[int]=None, **kwargs)Logs `values` to the current run.
class
src.accelerate.tracking.MLflowTrackerA `Tracker` class that supports `mlflow`.
method
src.accelerate.tracking.MLflowTracker.finish()End the active MLflow run.
method
src.accelerate.tracking.MLflowTracker.log(values:dict, step:Optional[int]=None, **kwargs)Logs `values` to the current run.
method
src.accelerate.tracking.MLflowTracker.log_figure(figure:Any, artifact_file:str, **save_kwargs)Logs an figure to the current run.
class
src.accelerate.tracking.SwanLabTrackerA `Tracker` class that supports `swanlab`.
method
src.accelerate.tracking.SwanLabTracker.finish()Closes `swanlab` writer
method
src.accelerate.tracking.SwanLabTracker.log(values:dict, step:Optional[int]=None, **kwargs)Logs `values` to the current run.
method
src.accelerate.tracking.SwanLabTracker.log_images(values:dict, step:Optional[int]=None, **kwargs)Logs `images` to the current run.
class
src.accelerate.tracking.TensorBoardTrackerA `Tracker` class that supports `tensorboard`.
method
src.accelerate.tracking.TensorBoardTracker.finish()Closes `TensorBoard` writer
method
src.accelerate.tracking.TensorBoardTracker.log(values:dict, step:Optional[int]=None, **kwargs)Logs `values` to the current run.
method
src.accelerate.tracking.TensorBoardTracker.log_images(values:dict, step:Optional[int]=None, **kwargs)Logs `images` to the current run.
class
src.accelerate.tracking.TrackioTrackerA `Tracker` class that supports `trackio`.
method
src.accelerate.tracking.TrackioTracker.finish()Closes `trackio` run
method
src.accelerate.tracking.TrackioTracker.log(values:dict, step:Optional[int]=None, **kwargs)Logs `values` to the current run.
class
src.accelerate.tracking.WandBTrackerA `Tracker` class that supports `wandb`.
method
src.accelerate.tracking.WandBTracker.finish()Closes `wandb` writer
method
src.accelerate.tracking.WandBTracker.log(values:dict, step:Optional[int]=None, **kwargs)Logs `values` to the current run.
method
src.accelerate.tracking.WandBTracker.log_images(values:dict, step:Optional[int]=None, **kwargs)Logs `images` to the current run.
class
src.accelerate.utils.dataclasses.ComputeEnvironmentRepresents a type of the compute environment.
class
src.accelerate.utils.dataclasses.DeepSpeedPluginThis plugin is used to integrate DeepSpeed.
class
src.accelerate.utils.dataclasses.DistributedTypeRepresents a type of distributed environment.
class
src.accelerate.utils.dataclasses.FP8BackendTypeRepresents the backend used for FP8.
class
src.accelerate.utils.dataclasses.FP8RecipeKwargsDeprecated.
class
src.accelerate.utils.dataclasses.SageMakerDistributedTypeRepresents a type of distributed environment.
class
src.accelerate.utils.deepspeed.DeepSpeedOptimizerWrapperInternal wrapper around a deepspeed optimizer.
class
src.accelerate.utils.deepspeed.DeepSpeedSchedulerWrapperInternal wrapper around a deepspeed scheduler.
func
src.accelerate.utils.environment.set_numa_affinity(local_process_index:int, verbose:Optional[bool]=None) -> NoneAssigns the current process to a specific NUMA node.
func
src.accelerate.utils.fsdp_utils.fsdp2_prepare_model(accelerator, model:torch.nn.Module) -> torch.nn.ModulePrepares the model for FSDP2 in-place.
func
src.accelerate.utils.imports.is_fp16_available()Checks if fp16 is supported
func
src.accelerate.utils.imports.is_fp8_available()Checks if fp8 is supported
func
src.accelerate.utils.launch.setup_fp8_env(args:argparse.Namespace, current_env:dict[str, str])Setup the FP8 environment variables.
class
src.accelerate.utils.megatron_lm.AbstractTrainStepAbstract class for batching, forward pass and loss handler.
class
src.accelerate.utils.megatron_lm.BertTrainStepBert train step class.
func
src.accelerate.utils.megatron_lm.BertTrainStep.forward_step(data_iterator, model)Forward step.
func
src.accelerate.utils.megatron_lm.BertTrainStep.get_batch_megatron(data_iterator)Build the batch.
func
src.accelerate.utils.megatron_lm.BertTrainStep.get_batch_transformer(data_iterator)Build the batch.
class
src.accelerate.utils.megatron_lm.GPTTrainStepGPT train step class.
func
src.accelerate.utils.megatron_lm.GPTTrainStep.forward_step(data_iterator, model)Forward step.
func
src.accelerate.utils.megatron_lm.GPTTrainStep.get_batch_megatron(data_iterator)Generate a batch
class
src.accelerate.utils.megatron_lm.T5TrainStepT5 train step class.
func
src.accelerate.utils.megatron_lm.T5TrainStep.forward_step(data_iterator, model)Forward step.
func
src.accelerate.utils.megatron_lm.T5TrainStep.get_batch_megatron(data_iterator)Build the batch.
func
src.accelerate.utils.megatron_lm.T5TrainStep.get_batch_transformer(data_iterator)Build the batch.
func
src.accelerate.utils.modeling.find_tied_parameters(model:torch.nn.Module, **kwargs) -> list[list[str]]Find the tied parameters in a given model.
func
src.accelerate.utils.modeling.id_tensor_storage(tensor:torch.Tensor) -> tuple[torch.device, int, int]Unique identifier to a tensor storage.
class
src.accelerate.utils.offload.PrefixedDatasetWill access keys in a given dataset by adding a prefix.
class
src.accelerate.utils.operations.DistributedOperationExceptionAn exception class for distributed operations.
func
src.accelerate.utils.other.get_free_port() -> intGets a free port on `localhost`.
func
src.accelerate.utils.other.get_pretty_name(obj)Gets a pretty name from `obj`.
func
src.accelerate.utils.other.has_repeated_blocks(module:torch.nn.Module) -> boolCheck whether the module has repeated blocks, i.e.
func
src.accelerate.utils.other.is_compiled_module(module:torch.nn.Module) -> boolCheck whether the module was compiled with torch.compile()
func
src.accelerate.utils.other.is_port_in_use(port:Optional[int]=None) -> boolChecks if a port is in use on `localhost`.
func
src.accelerate.utils.other.is_repeated_blocks(module:torch.nn.Module) -> boolCheck whether the module is a repeated block, i.e.
func
src.accelerate.utils.other.model_has_dtensor(model:torch.nn.Module) -> boolCheck if the model has DTensor parameters.
func
src.accelerate.utils.other.recursive_getattr(obj, attr:str)Recursive `getattr`.
func
src.accelerate.utils.other.save(obj, f, save_on_each_node:bool=False, safe_serialization:bool=False)Save the data to disk.
About this data
These signatures were extracted from the public source of huggingface/accelerate
using Python's ast module. Argument names, default values,
type annotations and return types are taken verbatim from the code.
Implementation bodies are never stored. See
how it works for details.