sdkagent

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

KindCount
Classes27
Functions26
Methods28

API list

methodsrc.accelerate.accelerator.Accelerator.context_parallel_rank() -> int
Context parallelism is not supported yet.
methodsrc.accelerate.accelerator.Accelerator.is_main_process()
True for one process only.
methodsrc.accelerate.accelerator.Accelerator.lomo_backward(loss:torch.Tensor, learning_rate:float) -> None
Runs backward pass on LOMO optimizers.
methodsrc.accelerate.accelerator.Accelerator.pipeline_parallel_rank() -> int
Pipeline parallelism is not supported yet.
methodsrc.accelerate.accelerator.Accelerator.tensor_parallel_rank() -> int
Returns the local rank for tensor parallelism.
classsrc.accelerate.commands.menu.input.KeyHandler
Metaclass that adds the key handlers to the class
funcsrc.accelerate.commands.menu.keymap.get_raw_chars()
Gets raw characters from inputs
funcsrc.accelerate.hooks.add_hook_to_module(module:nn.Module, hook:ModelHook, append:bool=False)
Adds a hook to a given module.
classsrc.accelerate.local_sgd.LocalSGD
A helper class to support local SGD on top of Accelerator.
classsrc.accelerate.logging.MultiProcessAdapter
An adapter to assist with logging in multiprocess.
classsrc.accelerate.optimizer.AcceleratedOptimizer
Internal wrapper around a torch optimizer.
methodsrc.accelerate.optimizer.AcceleratedOptimizer.eval()
Sets the optimizer to "eval" mode.
classsrc.accelerate.tracking.AimTracker
A `Tracker` class that supports `aim`.
methodsrc.accelerate.tracking.AimTracker.finish()
Closes `aim` writer
methodsrc.accelerate.tracking.AimTracker.log(values:dict, step:Optional[int]=None, **kwargs)
Logs `values` to the current run.
classsrc.accelerate.tracking.ClearMLTracker
A `Tracker` class that supports `clearml`.
methodsrc.accelerate.tracking.ClearMLTracker.finish()
Close the ClearML task.
methodsrc.accelerate.tracking.ClearMLTracker.log_images(values:dict, step:Optional[int]=None, **kwargs)
Logs `images` to the current run.
classsrc.accelerate.tracking.CometMLTracker
A `Tracker` class that supports `comet_ml`.
methodsrc.accelerate.tracking.CometMLTracker.finish()
Flush `comet-ml` writer
methodsrc.accelerate.tracking.CometMLTracker.log(values:dict, step:Optional[int]=None, **kwargs)
Logs `values` to the current run.
classsrc.accelerate.tracking.DVCLiveTracker
A `Tracker` class that supports `dvclive`.
methodsrc.accelerate.tracking.DVCLiveTracker.finish()
Closes `dvclive.Live()`.
methodsrc.accelerate.tracking.DVCLiveTracker.log(values:dict, step:Optional[int]=None, **kwargs)
Logs `values` to the current run.
classsrc.accelerate.tracking.MLflowTracker
A `Tracker` class that supports `mlflow`.
methodsrc.accelerate.tracking.MLflowTracker.finish()
End the active MLflow run.
methodsrc.accelerate.tracking.MLflowTracker.log(values:dict, step:Optional[int]=None, **kwargs)
Logs `values` to the current run.
methodsrc.accelerate.tracking.MLflowTracker.log_figure(figure:Any, artifact_file:str, **save_kwargs)
Logs an figure to the current run.
classsrc.accelerate.tracking.SwanLabTracker
A `Tracker` class that supports `swanlab`.
methodsrc.accelerate.tracking.SwanLabTracker.finish()
Closes `swanlab` writer
methodsrc.accelerate.tracking.SwanLabTracker.log(values:dict, step:Optional[int]=None, **kwargs)
Logs `values` to the current run.
methodsrc.accelerate.tracking.SwanLabTracker.log_images(values:dict, step:Optional[int]=None, **kwargs)
Logs `images` to the current run.
classsrc.accelerate.tracking.TensorBoardTracker
A `Tracker` class that supports `tensorboard`.
methodsrc.accelerate.tracking.TensorBoardTracker.finish()
Closes `TensorBoard` writer
methodsrc.accelerate.tracking.TensorBoardTracker.log(values:dict, step:Optional[int]=None, **kwargs)
Logs `values` to the current run.
methodsrc.accelerate.tracking.TensorBoardTracker.log_images(values:dict, step:Optional[int]=None, **kwargs)
Logs `images` to the current run.
classsrc.accelerate.tracking.TrackioTracker
A `Tracker` class that supports `trackio`.
methodsrc.accelerate.tracking.TrackioTracker.finish()
Closes `trackio` run
methodsrc.accelerate.tracking.TrackioTracker.log(values:dict, step:Optional[int]=None, **kwargs)
Logs `values` to the current run.
classsrc.accelerate.tracking.WandBTracker
A `Tracker` class that supports `wandb`.
methodsrc.accelerate.tracking.WandBTracker.finish()
Closes `wandb` writer
methodsrc.accelerate.tracking.WandBTracker.log(values:dict, step:Optional[int]=None, **kwargs)
Logs `values` to the current run.
methodsrc.accelerate.tracking.WandBTracker.log_images(values:dict, step:Optional[int]=None, **kwargs)
Logs `images` to the current run.
classsrc.accelerate.utils.dataclasses.ComputeEnvironment
Represents a type of the compute environment.
classsrc.accelerate.utils.dataclasses.DeepSpeedPlugin
This plugin is used to integrate DeepSpeed.
classsrc.accelerate.utils.dataclasses.DistributedType
Represents a type of distributed environment.
classsrc.accelerate.utils.dataclasses.FP8BackendType
Represents the backend used for FP8.
classsrc.accelerate.utils.dataclasses.FP8RecipeKwargs
Deprecated.
classsrc.accelerate.utils.dataclasses.SageMakerDistributedType
Represents a type of distributed environment.
classsrc.accelerate.utils.deepspeed.DeepSpeedOptimizerWrapper
Internal wrapper around a deepspeed optimizer.
classsrc.accelerate.utils.deepspeed.DeepSpeedSchedulerWrapper
Internal wrapper around a deepspeed scheduler.
funcsrc.accelerate.utils.environment.set_numa_affinity(local_process_index:int, verbose:Optional[bool]=None) -> None
Assigns the current process to a specific NUMA node.
funcsrc.accelerate.utils.fsdp_utils.fsdp2_prepare_model(accelerator, model:torch.nn.Module) -> torch.nn.Module
Prepares the model for FSDP2 in-place.
funcsrc.accelerate.utils.imports.is_fp16_available()
Checks if fp16 is supported
funcsrc.accelerate.utils.imports.is_fp8_available()
Checks if fp8 is supported
funcsrc.accelerate.utils.launch.setup_fp8_env(args:argparse.Namespace, current_env:dict[str, str])
Setup the FP8 environment variables.
classsrc.accelerate.utils.megatron_lm.AbstractTrainStep
Abstract class for batching, forward pass and loss handler.
classsrc.accelerate.utils.megatron_lm.BertTrainStep
Bert train step class.
funcsrc.accelerate.utils.megatron_lm.BertTrainStep.forward_step(data_iterator, model)
Forward step.
funcsrc.accelerate.utils.megatron_lm.BertTrainStep.get_batch_megatron(data_iterator)
Build the batch.
funcsrc.accelerate.utils.megatron_lm.BertTrainStep.get_batch_transformer(data_iterator)
Build the batch.
classsrc.accelerate.utils.megatron_lm.GPTTrainStep
GPT train step class.
funcsrc.accelerate.utils.megatron_lm.GPTTrainStep.forward_step(data_iterator, model)
Forward step.
funcsrc.accelerate.utils.megatron_lm.GPTTrainStep.get_batch_megatron(data_iterator)
Generate a batch
classsrc.accelerate.utils.megatron_lm.T5TrainStep
T5 train step class.
funcsrc.accelerate.utils.megatron_lm.T5TrainStep.forward_step(data_iterator, model)
Forward step.
funcsrc.accelerate.utils.megatron_lm.T5TrainStep.get_batch_megatron(data_iterator)
Build the batch.
funcsrc.accelerate.utils.megatron_lm.T5TrainStep.get_batch_transformer(data_iterator)
Build the batch.
funcsrc.accelerate.utils.modeling.find_tied_parameters(model:torch.nn.Module, **kwargs) -> list[list[str]]
Find the tied parameters in a given model.
funcsrc.accelerate.utils.modeling.id_tensor_storage(tensor:torch.Tensor) -> tuple[torch.device, int, int]
Unique identifier to a tensor storage.
classsrc.accelerate.utils.offload.PrefixedDataset
Will access keys in a given dataset by adding a prefix.
classsrc.accelerate.utils.operations.DistributedOperationException
An exception class for distributed operations.
funcsrc.accelerate.utils.other.get_free_port() -> int
Gets a free port on `localhost`.
funcsrc.accelerate.utils.other.get_pretty_name(obj)
Gets a pretty name from `obj`.
funcsrc.accelerate.utils.other.has_repeated_blocks(module:torch.nn.Module) -> bool
Check whether the module has repeated blocks, i.e.
funcsrc.accelerate.utils.other.is_compiled_module(module:torch.nn.Module) -> bool
Check whether the module was compiled with torch.compile()
funcsrc.accelerate.utils.other.is_port_in_use(port:Optional[int]=None) -> bool
Checks if a port is in use on `localhost`.
funcsrc.accelerate.utils.other.is_repeated_blocks(module:torch.nn.Module) -> bool
Check whether the module is a repeated block, i.e.
funcsrc.accelerate.utils.other.model_has_dtensor(model:torch.nn.Module) -> bool
Check if the model has DTensor parameters.
funcsrc.accelerate.utils.other.recursive_getattr(obj, attr:str)
Recursive `getattr`.
funcsrc.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.

Back to all 805 libraries