Spec-Zone.ru › PyTorch 2.14

Распределённые контрольные точки — torch.distributed.checkpoint

Создано: 16 нояб. 2022 г. | Последнее обновление: 08 июл. 2026 г.

Distributed Checkpoint (DCP) поддерживает параллельную загрузку и сохранение моделей из нескольких рангов. Он обрабатывает перераспределение фрагментов во время загрузки, что позволяет сохранять данные в одной топологии кластера, а загружать — в другой.

DCP отличается от torch.save и torch.load несколькими существенными особенностями:

  • Для каждой контрольной точки создаётся несколько файлов — как минимум по одному на ранг.
  • Операции выполняются на месте: это означает, что модель должна сначала выделить память для данных, а DCP использует уже выделенную память.

Для загрузки и сохранения контрольной точки используются следующие точки входа:

Дополнительные ресурсы:

  • Начало работы с Distributed Checkpoint (DCP)
  • Асинхронное сохранение с помощью Distributed Checkpoint (DCP)
  • Документация TorchTitan по контрольным точкам
  • Реализация DCP в TorchTitan
torch.distributed.checkpoint.optimizer.load_sharded_optimizer_state_dict(model_state_dict, optimizer_key, storage_reader, planner=None) [исходный код]

Загрузить state_dict совместно с фрагментированным состоянием оптимизатора FSDP.

В настоящее время это рекомендуемый способ создания контрольных точек FSDP.

Примеры:

>>> import torch.distributed.checkpoint as dist_cp
>>> # Save
>>> model: torch.nn.Model
>>> optim_params = model.parameters()
>>> optim = torch.optim.SGD(optim_params, lr=0.01)
>>> # Save
>>> with FSDP.state_dict_type(model, StateDictType.SHARDED_STATE_DICT):
>>>     state_dict = {
>>>         "optimizer": FSDP.optim_state_dict(model, optim),
>>>         "model": model.state_dict()
>>>     }
>>>     dist_cp.save_state_dict(
>>>         state_dict=optim_state,
>>>         storage_writer=dist_cp.FileSystemWriter("checkpoint"),
>>>         planner=dist_cp.DefaultSavePlanner(),
>>>     )
>>>
>>> # Load
>>> with FSDP.state_dict_type(model_tp, StateDictType.SHARDED_STATE_DICT):
>>>     model_state_dict = model_tp.state_dict()
>>>     checkpoint = {
>>>         "model": model_state_dict
>>>     }
>>>     dist_cp.load_state_dict(
>>>         state_dict=checkpoint,
>>>         storage_reader=dist_cp.FileSystemReader(checkpoint_file),
>>>         planner=dist_cp.DefaultLoadPlanner(),
>>>     )
>>>     model.load_state_dict(checkpoint["model_state"])
>>>
>>>     optim_state = dist_cp.load_sharded_optimizer_state_dict(
>>>         model_state_dict,
>>>         optimizer_key="optimizer",
>>>         storage_reader=dist_cp.FileSystemReader("checkpoint"),
>>>     )
>>>
>>>     flattened_osd = FSDP.optim_state_dict_to_load(
>>>        model, optim, optim_state["optimizer"]
>>>     )
>>>
>>>     optim.load_state_dict(flattened_osd)
Тип возвращаемого значения:

dict[str, StatefulT | Any]

torch.distributed.checkpoint.planner_helpers.create_read_items_for_chunk_list(fqn, checkpoint_md, local_chunks) [исходный код]

Создать список ReadItem на основе контрольной точки и локальных фрагментов.

Применяет алгоритм перераспределения фрагментов и вычисляет операции чтения, необходимые для удовлетворения local_chunks данными из контрольной точки, описанной в checkpoint_md.

Параметры:
  • fqn (str) – FQN элемента state_dict для передачи в ReadItem.
  • checkpoint_md (TensorStorageMetadata) – метаданные заданного тензора из контрольной точки.
  • local_chunks (List[ChunkStorageMetadata]) – локальные фрагменты, которые необходимо загрузить.
Возвращает:

Список ReadItem, которые удовлетворят требованиям всех входных фрагментов.

Тип возвращаемого значения:

list[ReadItem]

torch.distributed.checkpoint.default_planner.create_default_global_load_plan(all_plans) [исходный код]

Создать глобальный план загрузки, используемый DefaultLoadPlanner.

По умолчанию загрузка не требует глобальной координации, и в настоящее время эта функция не изменяет локальные планы.

Тип возвращаемого значения:

list[LoadPlan]

torch.distributed.checkpoint.default_planner.create_default_global_save_plan(all_plans, rewrite_index_hints=True) [исходный код]

Создать глобальный план и метаданные, используемые DefaultSavePlanner.

Метаданные формируются путём объединения метаданных всех WriteItem из предоставленных планов.

Единственное изменение при глобальном планировании — обновление подсказок индекса во всех объектах MetadataIndex, если rewrite_index_hints равно True.

Тип возвращаемого значения:

tuple[list[SavePlan], Metadata]

torch.distributed.checkpoint.default_planner.create_default_local_save_plan(state_dict, is_coordinator) [исходный код]

Создать SavePlan, используемый DefaultSavePlanner.

На рангах, не являющихся координатором, эта функция игнорирует тензоры и объекты, не являющиеся тензорами, и создаёт операции записи только для объектов ShardedTensor.

На ранге координатора создаются операции записи для всех значений.

Тип возвращаемого значения:

SavePlan

class torch.distributed.checkpoint.state_dict_saver.AsyncCheckpointerType(value) [исходный код]

Перечисление типов асинхронного средства создания контрольных точек.

class torch.distributed.checkpoint.state_dict_saver.AsyncSaveResponse(staging_completion, upload_completion) [исходный код]

Этот класс содержит объекты Future для завершения подготовки данных и выгрузки. Он возвращается функцией async_save(). staging_completion — это объект Future, указывающий, когда будет завершено локальное копирование state_dict. upload_completion — это объект Future, указывающий, когда будет завершено сохранение контрольной точки.

torch.distributed.checkpoint.state_dict_saver.save(state_dict, *, checkpoint_id=None, storage_writer=None, planner=None, process_group=None, no_dist=False, use_collectives=True) [исходный код]

Сохранить распределённую модель в стиле SPMD.

Эта функция отличается от torch.save() тем, что обрабатывает ShardedTensor и DTensor: каждый ранг сохраняет только свои локальные фрагменты.

Для каждого объекта Stateful (содержащего и state_dict, и load_state_dict) функция save перед сериализацией вызовет state_dict.

Предупреждение

Обратная совместимость сохранённых state_dict между версиями PyTorch не гарантируется.

Предупреждение

При использовании аргумента process_group убедитесь, что функцию save_state_dict вызывают только его ранги и что все данные в state_dict принадлежат ему.

Примечание

При сохранении контрольной точки для ShardingStrategy.HYBRID_SHARD FSDP функцию save_state_dict должна вызывать только одна группа фрагментов, которой также необходимо передать соответствующую группу процессов.

Примечание

Если группа процессов недоступна, эта функция предполагает, что требуется сохранить

state_dict в локальном процессе.

Параметры:
  • state_dict (Dict[str, Any]) – state_dict для сохранения.
  • checkpoint_id (Union[str, os.PathLike, None]) – Идентификатор этого экземпляра контрольной точки. Значение checkpoint_id зависит от хранилища. Это может быть путь к папке или файлу. Если используется хранилище «ключ — значение», это также может быть ключ. (По умолчанию: None)
  • storage_writer (Optional[StorageWriter]) – Экземпляр StorageWriter, используемый для выполнения операций записи. Если он не указан, DCP автоматически определит средство записи на основе checkpoint_id. Если checkpoint_id также равен None, будет вызвано исключение. (По умолчанию: None)
  • planner (Optional[SavePlanner]) – Экземпляр SavePlanner. Если он не указан, будет использоваться планировщик по умолчанию. (По умолчанию: None)
  • process_group (Optional[ProcessGroup]) – ProcessGroup для синхронизации между рангами. (По умолчанию: None)
  • no_dist (bool) – Если True, функция предполагает, что требуется загрузить контрольную точку на одном ранге/в одном процессе. (По умолчанию: False)
  • use_collectives (bool) – Если False, функция предполагает, что требуется сохранить контрольную точку без синхронизации между рангами. (По умолчанию: True) Эта конфигурация экспериментальная, поэтому используйте её с осторожностью. Она изменит формат сохранённой контрольной точки, которая может оказаться несовместимой с предыдущими версиями.
Возвращает:

Объект Metadata для сохранённой контрольной точки.

Тип возвращаемого значения:

Metadata

Пример

>>> my_model = MyModule()
>>> state_dict = {"model": my_model}
>>> fs_storage_writer = torch.distributed.checkpoint.FileSystemWriter(
...     "/checkpoint/1"
... )
>>> torch.distributed.checkpoint.save(
>>>     state_dict=state_dict,
>>>     storage_writer=fs_storage_writer,
>>> )

Примечание

save_state_dict использует коллективные операции для координации записи между рангами. Для групп процессов на основе NCCL внутренние тензорные представления объектов необходимо переместить на устройство GPU до начала обмена данными. В этом случае используемое устройство задаётся параметром torch.cuda.current_device(). Пользователь должен убедиться, что он настроен так, чтобы каждому рангу соответствовал отдельный GPU; для этого используется torch.cuda.set_device().

torch.distributed.checkpoint.state_dict_saver.async_save(state_dict, *, checkpoint_id=None, storage_writer=None, planner=None, process_group=None, async_checkpointer_type=AsyncCheckpointerType.THREAD, async_stager=None, no_dist=False, use_collectives=True) [исходный код]

Асинхронная версия save. Сначала этот код переносит state_dict из промежуточного хранилища в хранилище подготовки данных (по умолчанию — память CPU), а затем в отдельном потоке вызывает save.

Предупреждение

Эта функция экспериментальная и может измениться. Если вы предоставили async_stager, после сохранения последней контрольной точки вызовите close(). Внутренне созданные подготовители по умолчанию закрываются автоматически.

Параметры:
  • state_dict (Dict[str, Any]) – state_dict для сохранения.
  • checkpoint_id (Union[str, os.PathLike, None]) – Идентификатор этого экземпляра контрольной точки. Значение checkpoint_id зависит от хранилища. Это может быть путь к папке или файлу. Если используется хранилище «ключ — значение», это также может быть ключ. (По умолчанию: None)
  • storage_writer (Optional[StorageWriter]) – Экземпляр StorageWriter, используемый для выполнения этапов подготовки и сохранения. Если он не указан, DCP автоматически определит средство записи на основе checkpoint_id. Если checkpoint_id также равен None, будет вызвано исключение. (По умолчанию: None)
  • planner (Optional[SavePlanner]) – Экземпляр SavePlanner. Если он не указан, будет использоваться планировщик по умолчанию. (По умолчанию: None)
  • process_group (Optional[ProcessGroup]) – ProcessGroup для синхронизации между рангами. (По умолчанию: None)
  • async_checkpointer_type (AsyncCheckpointerType) – определяет, будет ли создание контрольной точки выполняться в отдельном потоке или процессе (по умолчанию: AsyncCheckpointerType.THREAD)
  • async_stager (AsyncStager) – предоставляет реализацию подготовки данных. Если storage_writer реализует AsyncStager, а async_stager не предоставлен, для подготовки данных будет использоваться storage_writer. Подготовители, предоставленные пользователем, остаются под его управлением, и их следует закрыть после сохранения последней контрольной точки.
  • no_dist (bool) – Если True, функция предполагает, что требуется сохранить контрольную точку на одном ранге/в одном процессе. (По умолчанию: False)
  • use_collectives (bool) – Если False, сохранить контрольную точку без координации между рангами. (По умолчанию: True) Эта конфигурация экспериментальная, поэтому используйте её с осторожностью. Она изменит формат сохранённой контрольной точки, которая может оказаться несовместимой с предыдущими версиями.
Возвращает:

Объект Future, содержащий результирующий объект Metadata из save.

Тип возвращаемого значения:

Future

Пример

>>> my_model = MyModule()
>>> state_dict = {"model": my_model}
>>> fs_storage_writer = torch.distributed.checkpoint.FileSystemWriter(
...     "/checkpoint/1"
... )
>>> checkpoint_future = torch.distributed.checkpoint.async_save(
>>>     state_dict=state_dict,
>>>     storage_writer=fs_storage_writer,
>>> )
>>>
>>> # ... do some work ...
>>>
>>> checkpoint_future.result()
torch.distributed.checkpoint.state_dict_saver.save_state_dict(state_dict, storage_writer, process_group=None, coordinator_rank=0, no_dist=False, planner=None) [исходный код]

Этот метод устарел. Используйте вместо него «save».

Тип возвращаемого значения:

Metadata

torch.distributed.checkpoint.state_dict_loader.load(state_dict, *, checkpoint_id=None, storage_reader=None, planner=None, process_group=None, no_dist=False) [исходный код]

Загрузить контрольную точку в распределённый state_dict в стиле SPMD.

Для этой функции API каждый ранг должен содержать одинаковые ключи в переданном state_dict. Несовпадение ключей может привести к зависанию или ошибкам. Если вы не уверены, можно использовать API utils._assert_same_keys для проверки (это может повлечь затраты на обмен данными).

Каждый ранг будет считывать минимально необходимый объём данных для выполнения запроса state_dict. При загрузке экземпляров ShardedTensor или DTensor каждый ранг считывает данные только для своих локальных фрагментов.

Для каждого объекта Stateful (содержащего и state_dict, и load_state_dict) функция load сначала вызовет state_dict перед десериализацией, а после её завершения — load_state_dict. Для каждого объекта, не являющегося Stateful, функция load десериализует объект, а затем заменит его в state_dict десериализованным объектом.

Предупреждение

Все тензоры в state_dict необходимо разместить на целевом устройстве до вызова этой функции.

Все данные, не являющиеся тензорами, загружаются с помощью torch.load() и изменяются на месте в state_dict.

Предупреждение

Пользователям необходимо вызвать load_state_dict для корневого модуля, чтобы обеспечить корректное выполнение последующей обработки при загрузке и распространение данных, не являющихся тензорами.

Параметры:
  • state_dict (Dict[str, Any]) – state_dict, в который необходимо загрузить контрольную точку.
  • checkpoint_id (Union[str, os.PathLike, None]) – Идентификатор этого экземпляра контрольной точки. Значение checkpoint_id зависит от хранилища. Это может быть путь к папке или файлу. Если используется хранилище «ключ — значение», это также может быть ключ. (По умолчанию: None)
  • storage_reader (Optional[StorageReader]) – Экземпляр StorageWriter, используемый для чтения. Если он не указан, DCP автоматически определит средство чтения на основе checkpoint_id. Если checkpoint_id также равен None, будет вызвано исключение. (По умолчанию: None)
  • planner (Optional[LoadPlanner]) – Экземпляр LoadPlanner. Если он не указан, будет использоваться планировщик по умолчанию. (По умолчанию: None)
  • process_group (Optional[ProcessGroup]) – ProcessGroup для синхронизации между рангами. (По умолчанию: None)
  • no_dist (bool) – Если True, функция предполагает, что требуется загрузить контрольную точку без синхронизации между рангами. (По умолчанию: False)
Возвращает:

None.

Тип возвращаемого значения:

None

Примеры
>>> my_model = MyModule()
>>> optimizer = Adagrad(my_model.parameters())
>>> model_state_dict = my_model.state_dict()
>>> fs_storage_reader = torch.distributed.checkpoint.FileSystemReader(
...     "/checkpoint/1"
... )
>>> torch.distributed.checkpoint.load_state_dict(
>>>     state_dict=model_state_dict,
>>>     storage_reader=fs_storage_reader,
>>> )
>>> # module.load_state_dict() function might have customized steps
>>> # to flush the state_dict, must call it to
>>> # ensure correct behavior.
>>> my_model.load_state_dict(model_state_dict)

Примечание

load_state_dict использует коллективные операции для координации чтения между рангами. Для групп процессов на основе NCCL внутренние тензорные представления объектов необходимо переместить на устройство GPU до начала обмена данными. В этом случае используемое устройство задаётся параметром torch.cuda.current_device(). Пользователь должен убедиться, что он настроен так, чтобы каждому рангу соответствовал отдельный GPU; для этого используется torch.cuda.set_device().

torch.distributed.checkpoint.state_dict_loader.load_state_dict(state_dict, storage_reader, process_group=None, coordinator_rank=0, no_dist=False, planner=None) [исходный код]

Этот метод устарел. Используйте вместо него «load».

Следующий модуль также полезен для дополнительной настройки механизмов подготовки данных, используемых при асинхронном создании контрольных точек (torch.distributed.checkpoint.async_save):

class torch.distributed.checkpoint.staging.AsyncStager(*args, **kwargs) [исходный код]

Этот протокол обеспечивает настройку и расширяемость dcp.async_save, позволяя пользователям задавать способ подготовки данных до параллельного выполнения обычного пути dcp.save. Ожидаемый порядок операций (конкретно определённый в torch.distributed.state_dict_saver.async_save) выглядит следующим образом:

  1. AsyncStager.stage_data(state_dict):

    Этот вызов предоставляет AsyncStager возможность подготовить state_dict. Цель подготовки в данном контексте — создать представление state_dict, безопасное для обучения: любые изменения данных модуля после завершения подготовки не должны отражаться в state_dict, возвращаемом этим методом. Например, в стандартном случае на CPU создаётся копия всего state_dict, что позволяет продолжать обучение, не рискуя изменить данные, которые сериализуются.

  2. Для state_dict, возвращённого stage, параллельно вызывается dcp.save. Этот вызов отвечает

    за сериализацию state_dict и его запись в хранилище.

  3. Если AsyncStager.should_synchronize_after_execute имеет значение True, этот метод будет вызван сразу после

    запуска потока сериализации и до возврата из dcp.async_save. Если задано значение False, предполагается, что пользователь определил собственную точку синхронизации для дальнейшей оптимизации задержки сохранения в цикле обучения (например, за счёт совмещения подготовки с прямым и обратным проходами). В этом случае пользователь должен вызвать AsyncStager.synchronize_staging в подходящий момент.

close() [исходный код]

Освободить все ресурсы, используемые подготовителем.

property should_synchronize_after_execute: bool

Следует ли выполнять синхронизацию после подготовки данных.

stage(state_dict) [исходный код]

Возвращает «подготовленную» копию state_dict. Предполагается, что подготовленная копия защищена от любых изменений, внесённых после завершения вызова stage.

Тип возвращаемого значения:

Future[dict[str, StatefulT | Any]] | dict[str, StatefulT | Any]

synchronize_staging() [исходный код]

Если stage выполняется асинхронно, этот метод необходимо вызвать, чтобы убедиться в завершении подготовки и безопасности изменения исходного state_dict.

class torch.distributed.checkpoint.staging.DefaultStager(config=StagingOptions(use_pinned_memory=True, use_shared_memory=True, use_async_staging=True, use_non_blocking_copy=True)) [исходный код]

DefaultStager предоставляет полнофункциональную реализацию подготовки, сочетающую несколько методов оптимизации для эффективной подготовки контрольных точек.

Процесс подготовки работает следующим образом: 1. Словарь состояния передаётся для подготовки (синхронно или асинхронно) 2. Тензоры копируются с GPU в оптимизированную память CPU 3. Операции CUDA синхронизируются, если используются неблокирующие копирования 4. Подготовленный словарь состояния возвращается или становится доступен через Future

Варианты использования:

# Synchronous staging stager = DefaultStager(StagingOptions(use_async_staging=False)) staged_dict = stager.stage(state_dict) stager.close()

# Asynchronous staging stager = DefaultStager(StagingOptions(use_async_staging=True)) future = stager.stage(state_dict) # … do other work … staged_dict = future.result() stager.close()

Особенности производительности:
  • Асинхронная подготовка обеспечивает наилучшую производительность, когда вычисления модели могут выполняться одновременно с операциями подготовки
  • Закреплённая память повышает скорость передачи данных между CPU и GPU, но требует больше памяти
  • Разделяемая память обеспечивает эффективный обмен данными с процессом создания контрольной точки
  • Неблокирующие копирования сокращают время простоя GPU при передаче данных в памяти
Потокобезопасность:

DefaultStager не является потокобезопасным. Каждый поток должен использовать собственный экземпляр либо должна быть обеспечена внешняя синхронизация.

close() [исходный код]

Освобождает все ресурсы, используемые DefaultStager. Завершает работу ThreadPoolExecutor, используемого для асинхронных операций подготовки, и очищает кэшированные хранилища базового StateDictStager. Этот метод следует вызывать, когда средство подготовки больше не требуется, чтобы избежать утечек ресурсов, особенно в долго работающих приложениях. После вызова close() средство подготовки нельзя использовать для дальнейших операций подготовки.

Пример использования:

stager = DefaultStager(StagingOptions(use_async_staging=True)) future = stager.stage(state_dict) result = future.result() stager.close() # Clean up all resources

stage(state_dict, **kwargs) [исходный код]

Эта функция отвечает за подготовку state_dict. Дополнительные сведения о подготовке см. в документации класса. Если use_async_staging имеет значение True, функция возвращает объект Future, который будет завершён по окончании подготовки. Если use_async_staging имеет значение False, функция возвращает полностью подготовленный state_dict.

Параметры:

state_dict (STATE_DICT_TYPE) – state_dict, который необходимо подготовить.

Тип возвращаемого значения:

dict[str, StatefulT | Any] | Future[dict[str, StatefulT | Any]]

synchronize_staging() [исходный код]

Если use_async_staging имеет значение True, этот метод ожидает завершения подготовки. Если use_async_staging имеет значение False, метод ничего не делает.

class torch.distributed.checkpoint.staging.StagingOptions(use_pinned_memory=True, use_shared_memory=True, use_async_staging=True, use_non_blocking_copy=True) [исходный код]

Параметры конфигурации поведения при подготовке контрольных точек.

Переменные:
  • use_pinned_memory (bool) – Включает выделение закреплённой памяти для ускорения передачи данных между CPU и GPU. Требуется доступность CUDA. По умолчанию: True
  • use_shared_memory (bool) – Включает разделяемую память для сценариев с несколькими процессами. Полезно, когда нескольким процессам необходим доступ к одним и тем же подготовленным данным. По умолчанию: True
  • use_async_staging (bool) – Включает асинхронную подготовку с использованием пула фоновых потоков. Позволяет выполнять вычисления одновременно с операциями подготовки. Требуется CUDA. По умолчанию: True
  • use_non_blocking_copy (bool) – Использует неблокирующее копирование памяти устройства с синхронизацией потоков. Повышает производительность, позволяя CPU продолжать работу во время передачи данных GPU. По умолчанию: True

Примечание

Функции, зависящие от CUDA, вызовут исключение, если CUDA недоступна.

class torch.distributed.checkpoint.staging.BlockingAsyncStager(cache_staged_state_dict=False, type_check=False) [исходный код]

Реализация AsyncStager, которая подготавливает state_dict в оперативной памяти CPU и блокирует выполнение до завершения копирования. Эта реализация также предоставляет возможность оптимизировать задержку подготовки с помощью закреплённой памяти.

Примечание: в этом случае synchronize_staging ничего не делает.

stage(state_dict) [исходный код]

Возвращает копию state_dict на CPU.

Тип возвращаемого значения:

dict[str, StatefulT | Any]

synchronize_staging() [исходный код]

Функция ничего не делает, поскольку подготовка выполняется блокирующим образом.

Помимо описанных выше точек входа, объекты Stateful, описание которых приведено ниже, позволяют дополнительно настраивать сохранение и загрузку.

class torch.distributed.checkpoint.stateful.Stateful(*args, **kwargs) [исходный код]

Протокол Stateful для объектов, которые можно сохранять в контрольную точку и восстанавливать из неё.

load_state_dict(state_dict) [исходный код]

Восстанавливает состояние объекта из предоставленного state_dict.

Параметры:

state_dict (dict[str, Any]) – Словарь состояния, из которого выполняется восстановление

state_dict() [исходный код]

Объекты должны возвращать представление своего состояния state_dict в виде словаря. Результат этой функции будет сохранён в контрольной точке, а затем восстановлен в load_state_dict().

Предупреждение

Из-за того, что восстановление контрольной точки выполняется на месте, эта функция также вызывается во время torch.distributed.checkpoint.load.

Возвращает:

Словарь состояния объекта

Тип возвращаемого значения:

Dict

В этом примере показано, как использовать PyTorch Distributed Checkpoint для сохранения модели FSDP.

Следующие типы определяют интерфейс ввода-вывода, используемый при работе с контрольными точками:

class torch.distributed.checkpoint.StorageReader [исходный код]

Интерфейс, используемый load_state_dict для чтения из хранилища.

В распределённой контрольной точке один экземпляр StorageReader выступает и координатором, и подчинённым узлом. При инициализации каждому экземпляру назначается его роль.

Подкласс должен ожидать следующую последовательность вызовов со стороны load_state_dict:

  1. (все ранги) устанавливает checkpoint_id, если пользователи передали допустимый checkpoint_id.
  2. (все ранги) вызывает read_metadata()
  3. (все ранги) вызывает set_up_storage_reader()
  4. (все ранги) вызывает prepare_local_plan()
  5. (координатор) вызывает prepare_global_plan()
  6. (все ранги) вызывает read_data()
abstract prepare_global_plan(plans) [исходный код]

Выполняет централизованное планирование загрузки из хранилища.

Этот метод вызывается только для экземпляра координатора.

Хотя этот метод может сформировать совершенно иной план, предпочтительный способ — сохранять специфичные для хранилища данные в LoadPlan::storage_data.

Параметры:

plans (list[LoadPlan]) – Список экземпляров LoadPlan, по одному для каждого ранга.

Возвращает:

Список преобразованных LoadPlan после глобального планирования хранилища

Тип возвращаемого значения:

list[LoadPlan]

abstract prepare_local_plan(plan) [исходный код]

Выполняет локальное планирование, специфичное для хранилища.

Хотя этот метод может сформировать совершенно иной план, рекомендуемый способ — сохранять специфичные для хранилища данные в LoadPlan::storage_data.

Параметры:

plan (LoadPlan) – Локальный план, предоставленный используемым LoadPlan.

Возвращает:

Преобразованный LoadPlan после локального планирования хранилища

Тип возвращаемого значения:

LoadPlan

abstract read_data(plan, planner) [исходный код]

Считывает все элементы из plan, используя planner для разрешения данных.

Подкласс должен вызывать LoadPlanner::load_bytes для десериализации объекта BytesIO в нужное место.

Подкласс должен вызывать LoadPlanner::resolve_tensor, чтобы получить доступ к тензорам, в которые необходимо загрузить данные.

StorageLayer отвечает за правильное планирование всех необходимых копирований между устройствами.

Параметры:
  • plan (LoadPlan) – Локальный план для выполнения
  • planner (LoadPlanner) – Объект планировщика, используемый для разрешения элементов.
Возвращает:

Future, который завершается после окончания всех операций чтения.

Тип возвращаемого значения:

Future[None]

abstract read_metadata(*args, **kwargs) [исходный код]

Считывает метаданные контрольной точки.

Возвращает:

Объект метаданных, связанный с загружаемой контрольной точкой.

Тип возвращаемого значения:

Metadata

abstract reset(checkpoint_id=None) [исходный код]

Вызов указывает на начало чтения новой контрольной точки. Если пользователи задали checkpoint_id для этого чтения контрольной точки, он может быть указан. Значение checkpoint_id зависит от хранилища. Это может быть путь к папке или файлу либо ключ для хранилища типа «ключ — значение».

Параметры:

checkpoint_id (Union[str, os.PathLike, None]) – Идентификатор этого экземпляра контрольной точки. Значение checkpoint_id зависит от хранилища. Это может быть путь к папке или файлу. Если хранилище представляет собой хранилище типа «ключ — значение», идентификатором может быть ключ. (По умолчанию: None)

abstract set_up_storage_reader(metadata, is_coordinator, *args, **kwargs) [исходный код]

Инициализирует этот экземпляр.

Параметры:
  • metadata (Metadata) – Используемая схема метаданных.
  • is_coordinator (bool) – Указывает, отвечает ли этот экземпляр за координацию контрольной точки.
abstract classmethod validate_checkpoint_id(checkpoint_id) [исходный код]

Проверяет, поддерживается ли переданный checkpoint_id хранилищем. Это позволяет автоматически выбирать хранилище.

Тип возвращаемого значения:

bool

class torch.distributed.checkpoint.StorageWriter [исходный код]

Интерфейс, используемый save_state_dict для записи в хранилище.

Один экземпляр StorageWriter выступает как координатор и как ведомый в распределенной контрольной точке. При инициализации каждому экземпляру назначается роль.

Подкласс должен ожидать следующую последовательность вызовов.

  1. (все ранги) задают checkpoint_id, если пользователи передали допустимый checkpoint_id.
  2. (все ранги) вызывают set_up_storage_writer()
  3. (все ранги) вызывают prepare_local_plan()
  4. (координатор) вызывает prepare_global_plan()
  5. (все ранги) вызывают write_data()
  6. (координатор) вызывает finish()
abstract finish(metadata, results) [исходный код]

Записывает метаданные и помечает текущую контрольную точку как успешно созданную.

Фактический формат/схема, используемые для сериализации metadata, являются деталями реализации. Единственное требование — возможность восстановить тот же граф объектов.

Параметры:
  • metadata (Metadata) – метаданные новой контрольной точки
  • results (list[list[WriteResult]]) – список объектов WriteResult со всех рангов.
Возвращает:

None

Тип возвращаемого значения:

None

abstract prepare_global_plan(plans) [исходный код]

Выполняет централизованное планирование хранения.

Этот метод вызывается только для экземпляра координатора.

Хотя этот метод может создавать совершенно другой план, предпочтительно сохранять специфичные для хранилища данные в SavePlan::storage_data.

Параметры:

plans (list[SavePlan]) – список экземпляров SavePlan, по одному для каждого ранга.

Возвращает:

Список преобразованных SavePlan после глобального планирования хранения

Тип возвращаемого значения:

list[SavePlan]

abstract prepare_local_plan(plan) [исходный код]

Выполняет локальное планирование с учетом особенностей хранилища.

Хотя этот метод может создавать совершенно другой план, рекомендуется сохранять специфичные для хранилища данные в SavePlan::storage_data.

Параметры:

plan (SavePlan) – локальный план используемого SavePlanner.

Возвращает:

Преобразованный SavePlan после локального планирования хранения

Тип возвращаемого значения:

SavePlan

abstract reset(checkpoint_id=None) [исходный код]

Вызов указывает на начало записи новой контрольной точки. checkpoint_id может быть указан, если пользователи задали checkpoint_id для этой записи контрольной точки. Значение checkpoint_id зависит от хранилища. Это может быть путь к папке/файлу или ключ в хранилище «ключ-значение».

Параметры:

checkpoint_id (Union[str, os.PathLike, None]) – идентификатор этого экземпляра контрольной точки. Значение checkpoint_id зависит от хранилища. Это может быть путь к папке или файлу. Если хранилище использует формат «ключ-значение», это также может быть ключ. (По умолчанию: None)

abstract set_up_storage_writer(is_coordinator, *args, **kwargs) [исходный код]

Инициализирует этот экземпляр.

Параметры:

is_coordinator (bool) – указывает, отвечает ли этот экземпляр за координацию контрольной точки.

storage_meta() [исходный код]

Возвращает метаданные, специфичные для хранилища. Они используются для сохранения в контрольной точке дополнительной информации, которая может быть полезна для наблюдаемости на уровне запросов. StorageMeta передается в SavePlanner при вызове сохранения. По умолчанию возвращается None.

TODO: привести пример

Тип возвращаемого значения:

StorageMeta | None

abstract classmethod validate_checkpoint_id(checkpoint_id) [исходный код]

Проверяет, поддерживается ли заданный checkpoint_id хранилищем. Это позволяет автоматически выбирать хранилище.

Тип возвращаемого значения:

bool

abstract write_data(plan, planner) [исходный код]

Записывает все элементы из plan, используя planner для получения данных.

Подкласс должен вызывать SavePlanner::resolve_data для каждого элемента плана, чтобы получить доступ к базовому объекту для записи.

Подклассам следует вызывать resolve_data по мере необходимости, поскольку это может выделять память. Для тензоров следует учитывать следующее:

  • Они могут находиться на любом устройстве, в том числе не совпадающем с устройством в WriteItem::tensor_data
  • Они могут быть представлениями или неконт contiguousными. Сохранять нужно только проекцию.
Параметры:
  • plan (SavePlan) – план сохранения для выполнения.
  • planner (SavePlanner) – объект планировщика, используемый для получения данных из элементов.
Возвращает:

объект future, который после завершения возвращает список WriteResult

Тип возвращаемого значения:

Future[list[WriteResult]]

Следующие типы определяют метаданные, используемые при работе с контрольными точками:

class torch.distributed.checkpoint.metadata.StorageMeta(checkpoint_id: str | os.PathLike | None = None, save_id: str | None = None, load_id: str | None = None, modules: list[str] = <factory>) [исходный код]
class torch.distributed.checkpoint.metadata.TensorProperties(dtype=<factory>, layout=torch.strided, requires_grad=False, memory_format=torch.contiguous_format, pin_memory=False) [исходный код]

Свойства, используемые для создания Tensor

class torch.distributed.checkpoint.metadata.TensorStorageMetadata(properties: torch.distributed.checkpoint.metadata.TensorProperties, size: torch.Size, chunks: list[torch.distributed.checkpoint.metadata.ChunkStorageMetadata]) [исходный код]

Следующие типы определяют интерфейс планировщика, используемый при работе с контрольными точками:

class torch.distributed.checkpoint.LoadPlanner [исходный код]

Абстрактный класс, определяющий протокол, который load_state_dict использует для планирования процесса загрузки.

LoadPlanner — это объекты с состоянием, позволяющие настраивать весь процесс загрузки.

LoadPlanner служит прокси для доступа к state_dict, поэтому любые его преобразования будут видны во всем процессе.

Подкласс планировщика может ожидать следующую последовательность вызовов при выполнении load_state_dict:

  1. set_up_planner — вызывается на всех рангах.

    Сигнализирует о начале загрузки контрольной точки.

  2. create_local_plan — вызывается на всех рангах.

    Обрабатывает state_dict и создает LoadPlan, который будет передан для глобального планирования.

  3. create_global_plan — вызывается только на ранге координатора.

    Получает LoadPlan со всех рангов и принимает глобальные решения.

  4. load_bytes — вызывается несколько раз на каждом ранге

    Вызывается один раз для каждого значения в state_dict, не являющегося тензором.

  5. resolve_tensor и commit_tensor — вызываются несколько раз на каждом ранге

    Вызываются парой для каждого значения Tensor в state_dict.

Рекомендуется наследоваться от DefaultLoadPlanner, а не напрямую реализовывать этот интерфейс, поскольку большинство изменений можно выразить изменением одного метода.

Существуют два распространенных способа расширения:

Преобразование state_dict. Это самый простой способ расширить процесс загрузки, поскольку он не требует понимания особенностей работы LoadPlan. Необходимо сохранить ссылку на исходный state_dict, так как загрузка выполняется на месте, и поэтому ее также нужно выполнять на месте

>>> class RenamePlanner(DefaultLoadPlanner):
>>>     def set_up_planner(
>>>         self,
>>>         state_dict: STATE_DICT_TYPE,
>>>         metadata: Metadata,
>>>         is_coordinator: bool,
>>>     ) -> None:
>>>         self.original_state_dict = state_dict
>>>         state_dict = {"foo_" + k: v for k, v in state_dict.items()}
>>>
>>>         if self.flatten_sharded_tensors:
>>>             state_dict = _flatten_sharded_tensors(state_dict)
>>>
>>>         if self.flatten_state_dict:
>>>             state_dict, self.mappings = flatten_state_dict(state_dict)
>>>
>>>         self.state_dict = state_dict
>>>         self.metadata = metadata
>>>         self.is_coordinator = is_coordinator
>>>
>>>     def load_bytes(self, read_item, value):
>>> # Remove the "foo_" prefix
>>>         self.original_state_dict[read_item.dest_index.fqn[4:]] = torch.load(value, weights_only=False)

Изменение resolve_tensor и commit_tensor для выполнения преобразования во время загрузки.

>>> class MetaModelMaterialize(DefaultSavePlanner):
>>>     def resolve_tensor(self, read_item):
>>>         tensor = super().resolve_tensor(read_item)
>>>         return torch.empty_like(tensor, device="cpu")
>>>
>>>     def commit_tensor(self, read_item, tensor):
>>>         self.state_dict[read_item.dest_index.fqn] = tensor
abstract commit_tensor(read_item, tensor) [исходный код]

Вызывается после того, как StorageReader завершит загрузку данных в tensor.

Переданный тензор — это тот же тензор, который был возвращен вызовом resolve_tensor. Этот метод нужен только в том случае, если LoadPlanner должен выполнить постобработку tensor перед копированием обратно в тензор из state_dict.

Содержимое тензора будет подчиняться модели синхронизации его устройства.

abstract create_global_plan(global_plan) [исходный код]

Вычисляет глобальный план загрузки и возвращает планы для каждого ранга.

. Примечание: вызывается только на ранге координатора

Тип возвращаемого значения:

list[LoadPlan]

abstract create_local_plan() [исходный код]

Создает LoadPlan на основе state_dict и метаданных, переданных в set_up_planner.

. Примечание: вызывается на каждом ранге.

Тип возвращаемого значения:

LoadPlan

abstract finish_plan(central_plan) [исходный код]

Принимает план от координатора и возвращает итоговый LoadPlan.

Тип возвращаемого значения:

LoadPlan

abstract load_bytes(read_item, value) [исходный код]

Загружает элемент, описанный в read_item``and ``value.

Этот метод должен изменять базовый state_dict на месте.

Содержимое value определяется SavePlanner, использованным для создания загружаемой контрольной точки.

resolve_bytes(read_item) [исходный код]

Возвращает BytesIO, который StorageReader будет использовать для загрузки read_item.

BytesIO должен ссылаться на тот же объект, что и в базовом state_dict, поскольку StorageReader заменит его содержимое.

Тип возвращаемого значения:

BytesIO

abstract resolve_tensor(read_item) [исходный код]

Возвращает тензор, описанный в read_item, который StorageReader будет использовать для загрузки read_item.

Тензор должен ссылаться на тот же объект, что и в базовом state_dict, поскольку StorageReader заменит его содержимое. Если по какой-либо причине это невозможно, планировщик может использовать метод commit_tensor, чтобы скопировать данные обратно в state_dict.

Тип возвращаемого значения:

Tensor

abstract set_up_planner(state_dict, metadata=None, is_coordinator=False) [исходный код]

Инициализирует этот экземпляр для загрузки данных в state_dict.

. Примечание: вызывается на каждом ранге.

class torch.distributed.checkpoint.LoadPlan(items: list[torch.distributed.checkpoint.planner.ReadItem], storage_data: Any = None, planner_data: Any = None) [исходный код]
class torch.distributed.checkpoint.ReadItem(type: torch.distributed.checkpoint.planner.LoadItemType, dest_index: torch.distributed.checkpoint.metadata.MetadataIndex, dest_offsets: torch.Size, storage_index: torch.distributed.checkpoint.metadata.MetadataIndex, storage_offsets: torch.Size, lengths: torch.Size) [исходный код]
class torch.distributed.checkpoint.SavePlanner [исходный код]

Абстрактный класс, определяющий протокол, который save_state_dict использует для планирования процесса сохранения.

SavePlanner — это объекты с состоянием, позволяющие настраивать весь процесс сохранения.

SavePlanner служит прокси для доступа к state_dict, поэтому любые его преобразования будут видны во всем процессе.

Подкласс планировщика может ожидать следующую последовательность вызовов при выполнении save_state_dict:

  1. set_up_planner — вызывается на всех рангах.

    Сигнализирует о начале сохранения контрольной точки.

  2. create_local_plan — вызывается на всех рангах.

    Обрабатывает state_dict и создает SavePlan, который будет передан для глобального планирования.

  3. create_global_plan — вызывается только на ранге координатора.

    Получает SavePlan со всех рангов и принимает глобальные решения.

  4. finish_plan — вызывается на всех рангах.

    Позволяет каждому рангу скорректировать план с учетом решений глобального планирования.

  5. resolve_data — вызывается несколько раз на каждом ранге

    Находит значение в state_dict для записи на уровне хранилища.

Рекомендуется наследоваться от DefaultSavePlanner, а не напрямую реализовывать этот интерфейс, поскольку большинство изменений можно выразить изменением одного метода.

Существуют 3 распространенных способа расширения:

Преобразование state_dict. Это самый простой способ расширить процесс сохранения, поскольку он не требует понимания особенностей работы SavePlan:

>>> class RenamePlanner(DefaultSavePlanner):
>>>     def set_up_planner(
>>>         self,
>>>         state_dict: STATE_DICT_TYPE,
>>>         storage_meta: Optional[StorageMeta],
>>>         is_coordinator: bool,
>>>     ) -> None:
>>> # prefix all keys with `foo_``
>>>         super().set_up_planner({"foo_" + k: v for k, v in state_dict.items()}, storage_meta, is_coordinator)

Совместное изменение локального плана и поиска данных. Это полезно, когда требуется точно контролировать способ сохранения данных

>>> class FP16Planner(DefaultSavePlanner):
>>>     def create_local_plan(self):
>>>         plan = super().create_local_plan()
>>>         for p in plan:
>>>             if p.tensor_data is not None:
>>>                 p.tensor_data.properties.dtype = torch.float16
>>>         return plan
>>>
>>>     def resolve_data(self, write_item):
>>>         item = super().resolve_data(write_item)
>>>         return item if write_item.type == WriteItemType.BYTE_IO else item.to(torch.float16)

Использование этапа глобального планирования для принятия централизованных решений, которые не могут быть приняты каждым рангом по отдельности

>>> from itertools import zip_longest
>>> from dataclasses import replace
>>> class DDPLoadBalancingPlanner(DefaultSavePlanner):
>>> # This uses the default local plan behavior of having all non-sharded writes in rank 0
>>> # This sample doesn't handle ShardedTensors
>>>     def create_global_plan(self, all_plans):
>>>         iters = [iter(all_plans[0].items)] * len(all_plans)
>>>         items_per_rank = [
>>>             [item for item in items if item is not None]
>>>             for items in zip(*zip_longest(*iters), strict=True)
>>>         ]
>>>         all_plans = [
>>>             replace(plan, items=items)
>>>             for plan, items in zip(all_plans, items_per_rank, strict=True)
>>>         ]
>>>         return super().create_global_plan(all_plans)

Наконец, некоторым планировщикам нужно сохранять в контрольной точке дополнительные метаданные. Для этого каждый ранг добавляет свои элементы данных в локальный план, а глобальный планировщик объединяет их:

>>> class SaveExtraDataPlanner(DefaultSavePlanner):
>>>     def create_local_plan(self) -> SavePlan:
>>>         plan = super().create_local_plan()
>>>         return replace(plan, planner_data="per-rank-data")
>>>
>>>     def create_global_plan(self, all_plans: List[SavePlan]) -> Tuple[List[SavePlan], Metadata]:
>>>         global_plan, metadata = super().create_global_plan(all_plans)
>>>         merged_data = [p.planner_data for p in global_plan]
>>>         metadata = replace(metadata, planner_data=merged_data)
>>>         return global_plan, metadata
abstract create_global_plan(all_plans) [исходный код]

Вычисляет глобальный план контрольной точки и возвращает локальный план каждого ранга.

Вызывается только на ранге координатора.

Тип возвращаемого значения:

tuple[list[SavePlan], Metadata]

abstract create_local_plan() [исходный код]

Вычисляет план сохранения для текущего ранга.

Этот план будет объединен и передан в create_global_plan. Специфичные для планировщика данные можно передать через SavePlan::planner_data.

Вызывается на всех рангах.

Тип возвращаемого значения:

SavePlan

abstract finish_plan(new_plan) [исходный код]

Объединяет план, созданный методом create_local_plan, с результатом create_global_plan.

Вызывается на всех рангах.

Тип возвращаемого значения:

SavePlan

abstract resolve_data(write_item) [исходный код]

Преобразует и подготавливает write_item из state_dict для хранения, обеспечивая идемпотентность и потокобезопасность.

Находит объект, связанный с write_item в state_dict, и применяет необходимые преобразования (например, сериализацию) до передачи объекта уровню хранилища.

Вызывается несколько раз на каждом ранге, как минимум один раз для каждого WriteItem в итоговом SavePlan.

Этот метод должен быть идемпотентным и потокобезопасным. Реализации StorageWriter могут вызывать его так часто, как требуется.

Любые преобразования, выделяющие память, следует выполнять лениво при вызове этого метода, чтобы уменьшить пиковый объем памяти, необходимый для создания контрольной точки.

Возвращаемые тензоры могут находиться на любом устройстве и иметь любой формат; они также могут быть представлениями. Уровень хранилища отвечает за выбор способа их сохранения.

Тип возвращаемого значения:

Tensor | BytesIO

abstract set_up_planner(state_dict, storage_meta=None, is_coordinator=False) [исходный код]

Инициализирует этот планировщик для сохранения state_dict.

Реализациям следует сохранить эти значения, поскольку в дальнейшем процессе сохранения они передаваться не будут.

Вызывается на всех рангах.

class torch.distributed.checkpoint.SavePlan(items: list[torch.distributed.checkpoint.planner.WriteItem], storage_data: Any = None, planner_data: Any = None, usable: bool = True) [исходный код]
class torch.distributed.checkpoint.planner.WriteItem(index, type, bytes_io_data=None, tensor_data=None) [исходный код]

Класс данных, содержащий сведения о том, что необходимо записать в хранилище.

tensor_storage_size() [исходный код]

Вычисляет размер базового тензора в хранилище или возвращает None, если это не запись тензора.

Возвращает:

Optional[int]: размер базового тензора в хранилище в байтах, если он есть.

Тип возвращаемого значения:

int | None

class torch.distributed.checkpoint.planner.BytesIOWriteData(nbytes: int) [исходный код]
class torch.distributed.checkpoint.planner.TensorWriteData(chunk: torch.distributed.checkpoint.metadata.ChunkStorageMetadata, properties: torch.distributed.checkpoint.metadata.TensorProperties, size: torch.Size) [исходный код]

Мы предоставляем уровень хранения на основе файловой системы:

class torch.distributed.checkpoint.filesystem.FileSystemBase [исходный код]
class torch.distributed.checkpoint.filesystem.FileSystem [исходный код]
class torch.distributed.checkpoint.filesystem.SerializationFormat(value) [исходный код]

Перечисление.

class torch.distributed.checkpoint.FileSystemReader(path, _extension_registry=None) [исходный код]
property checkpoint_id: str | PathLike

возвращает checkpoint_id, который будет использоваться для загрузки контрольной точки.

class torch.distributed.checkpoint.FileSystemWriter(path, single_file_per_rank=True, sync_files=True, thread_count=1, per_thread_copy_ahead=10000000, cache_staged_state_dict=False, overwrite=True, _extensions=None, serialization_format=SerializationFormat.TORCH_SAVE) [исходный код]

Базовая реализация StorageWriter с использованием файлового ввода-вывода.

Эта реализация предполагает и упрощает следующее:

  • Путь к контрольной точке указывает на пустой или несуществующий каталог.
  • Создание файла является атомарным

Контрольная точка состоит из одного файла на каждый запрос записи, а также глобального файла .metadata с сериализованными метаданными, если включена координация рангов. Если координация рангов НЕ включена, используется локальный для ранга файл __{rank}.metadata с сериализованными метаданными.

stage(state_dict) [исходный код]

Переопределение AsyncStager.stage

Тип возвращаемого значения:

dict[str, StatefulT | Any]

Мы также предоставляем другие уровни хранения, в том числе для работы с safetensors от HuggingFace:

.. autoclass:: torch.distributed.checkpoint.HuggingFaceStorageReader :members:

.. autoclass:: torch.distributed.checkpoint.HuggingFaceStorageWriter :members:

.. autoclass:: torch.distributed.checkpoint.QuantizedHuggingFaceStorageReader :members:

Мы предоставляем реализации по умолчанию для LoadPlanner и SavePlanner, которые поддерживают все конструкции torch.distributed, такие как FSDP, DDP, ShardedTensor и DistributedTensor.

class torch.distributed.checkpoint.DefaultSavePlanner(flatten_state_dict=True, flatten_sharded_tensors=True, dedup_replicated_tensors=None, dedup_save_to_lowest_rank=False, enable_plan_caching=False) [исходный код]
lookup_object(index) [исходный код]

Расширение интерфейса планировщика, упрощающее расширение планировщика по умолчанию.

Тип возвращаемого значения:

Any

transform_object(write_item, object) [исходный код]

Расширение интерфейса планировщика, упрощающее расширение планировщика по умолчанию.

class torch.distributed.checkpoint.DefaultLoadPlanner(flatten_state_dict=True, flatten_sharded_tensors=True, allow_partial_load=False) [исходный код]

DefaultLoadPlanner, добавляющий несколько возможностей к LoadPlanner.

В частности, он добавляет следующее:

flatten_state_dict: обработка state_dict с вложенными словарями flatten_sharded_tensors: для FSDP в режиме параллелизма 2D allow_partial_load: если False, возникнет ошибка времени выполнения, если ключ присутствует в state_dict, но отсутствует в контрольной точке.

lookup_tensor(index) [исходный код]

Расширение интерфейса планировщика, упрощающее расширение планировщика по умолчанию.

Тип возвращаемого значения:

Tensor

transform_tensor(read_item, tensor) [исходный код]

Расширение интерфейса планировщика, упрощающее расширение планировщика по умолчанию.

Из-за решений, принятых при проектировании устаревших API, словари состояния FSDP и DDP могут иметь разные ключи или полные имена (например, layer1.weight), даже если исходная модель без параллелизма идентична. Кроме того, FSDP предоставляет различные типы словарей состояния модели, например полные и сегментированные словари состояния. Также в словарях состояния оптимизатора для идентификации параметров используются идентификаторы параметров вместо полных имен, что может вызывать проблемы при использовании параллелизма (например, конвейерного параллелизма).

Для решения этих задач мы предлагаем набор API, позволяющих легко управлять state_dict. get_model_state_dict() возвращает словарь состояния модели с ключами, совпадающими с ключами словаря состояния модели без параллелизма. Аналогично, get_optimizer_state_dict() предоставляет словарь состояния оптимизатора с единообразными ключами для всех применённых видов параллелизма. Для обеспечения такой согласованности get_optimizer_state_dict() преобразует идентификаторы параметров в полные имена, совпадающие с именами в словаре состояния модели без параллелизма.

Обратите внимание, что результаты, возвращаемые этими API, можно напрямую использовать с методами torch.distributed.checkpoint.save() и torch.distributed.checkpoint.load() без каких-либо дополнительных преобразований.

Методы set_model_state_dict() и set_optimizer_state_dict() предназначены для загрузки state_dict модели и оптимизатора, сформированных соответствующими API получения.

Обратите внимание, что set_optimizer_state_dict() можно вызывать только до вызова backward() или после вызова step() для оптимизаторов.

Обратите внимание, что эта возможность является экспериментальной, и сигнатуры API могут измениться в будущем.

torch.distributed.checkpoint.state_dict.get_state_dict(model, optimizers, *, submodules=None, options=None) [исходный код]

Возвращает state_dict модели и state_dict оптимизаторов.

get_state_dict может обрабатывать любой модуль, параллелизованный с помощью PyTorch FSDP/fully_shard, DDP/replicate, tensor_parallel/parallelize_module, а также любые комбинации этих видов параллелизма. Основные функции get_state_dict: 1.) возвращение state_dict модели и оптимизатора, который можно переразбить на другое количество тренеров и/или с другими видами параллелизма. 2.) скрытие API state_dict, специфичных для параллелизма. Пользователям не нужно вызывать эти API. 3.) проверка корректности результирующего state_dict.

Ключи результирующего словаря состояния — канонические FQN (полные имена). Канонический FQN — это полное имя, основанное на положении параметра в иерархии nn.Module. Точнее, каноническим FQN параметра является FQN, возвращаемый module.named_parameters() или module.named_buffers(), если модуль не распределён с использованием параллелизма. Поскольку оптимизатор внутри использует идентификаторы параметров для их представления, при вызове этого API идентификаторы параметров преобразуются в канонические FQN.

get_state_dict также может обрабатывать модуль без параллелизма. В таком случае get_state_dict выполняет только одну функцию — преобразует идентификаторы параметров оптимизатора в канонические FQN.

Пример

>>> import torch
>>> from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
>>> from torch.nn.parallel import DistributedDataParallel as DDP
>>> from torch.distributed.checkpoint.state_dict import get_state_dict
>>> fsdp_model = FSDP(copy.deepcopy(model))
>>> fsdp_optim = torch.optim.Adam(model.parameters(), lr=1e-3)
>>> ddp_model = DDP(copy.deepcopy(model))
>>> ddp_optim = torch.optim.Adam(model.parameters(), lr=1e-3)
>>> ddp_state_dict, ddp_optim_state_dict = get_state_dict(ddp_model, ddp_optim)
>>> fsdp_state_dict, fsdp_optim_state_dict = get_state_dict(
...     fsdp_model, fsdp_optim
... )
>>> # if we simply call ddp_model.state_dict() and fsdp_model.state_dict(),
>>> # the asserts will fail.
>>> assert ddp_state_dict == fsdp_state_dict
>>> assert ddp_optim_state == fsdp_optim_state_dict
Параметры:
  • model (nn.Module) – nn.Module модели.
  • optimizers (Union[None, Optimizer, Iterable[Optimizer]]) – оптимизаторы, используемые для оптимизации model.
  • submodules (устарело) – Optional[set[nn.Module]]: возвращать только параметры модели, относящиеся к подмодулям.
  • options (StateDictOptions) – параметры, управляющие возвращаемыми model state_dict и optimizer state_dict. Подробности см. в StateDictOptions.
Возвращает:

Tuple, содержащий model state_dict и optimizer state_dict.

Тип возвращаемого значения:

Tuple[Dict[str, ValueType], OptimizerStateType]

torch.distributed.checkpoint.state_dict.get_model_state_dict(model, *, submodules=None, options=None) [исходный код]

Возвращает state_dict модели model.

Подробные сведения об использовании см. в get_state_dict.

Параметры:
  • model (nn.Module) – nn.Module модели.
  • submodules (устарело) – Optional[set[nn.Module]]: возвращать только параметры модели, относящиеся к подмодулям.
  • options (StateDictOptions) – параметры, управляющие возвращаемыми model state_dict и optimizer state_dict. Подробности см. в StateDictOptions.
Возвращает:

state_dict для model.

Тип возвращаемого значения:

Dict[str, ValueType]

torch.distributed.checkpoint.state_dict.get_optimizer_state_dict(model, optimizers, *, submodules=None, options=None) [исходный код]

Возвращает объединённый state_dict для оптимизаторов.

Подробные сведения об использовании см. в get_state_dict.

Параметры:
  • model (nn.Module) – nn.Module модели.
  • optimizers (Union[None, Optimizer, Iterable[Optimizer]]) – оптимизаторы, используемые для оптимизации model.
  • submodules (устарело) – Optional[set[nn.Module]]: возвращать только параметры модели, относящиеся к подмодулям.
  • options (StateDictOptions) – параметры, управляющие возвращаемыми model state_dict и optimizer state_dict. Подробности см. в StateDictOptions.
Возвращает:

state_dict для optimizers.

Тип возвращаемого значения:

OptimizerStateType

torch.distributed.checkpoint.state_dict.set_state_dict(model, optimizers, *, model_state_dict, optim_state_dict, options=None) [исходный код]

Загружает state_dict модели и state_dict оптимизаторов.

Метод, обратный get_state_dict, который устанавливает state_dict для модели и оптимизаторов. Переданные model_state_dict и optim_state_dict не обязательно должны быть возвращены методом get_state_dict, но должны отвечать следующим требованиям: 1) все FQN являются каноническими FQN, определёнными в get_state_dict; 2) если тензор разбит на части, он должен быть ShardedTensor или DTensor; 3) optimizer state_dict не может содержать идентификаторы параметров — в качестве ключей должны использоваться канонические FQN.

WARN: set_state_dict can only be called before backward() or after step()

вызывается для оптимизаторов. В противном случае состояния оптимизатора не будут корректно инициализированы.

Параметры:
  • model (nn.Module) – nn.Module модели.
  • optimizers (Union[Optimizer, Iterable[Optimizer]]) – оптимизаторы, используемые для оптимизации model.
  • model_state_dict (Dict[str, ValueType]) – (Union[Dict[nn.Module, Dict[str, ValueType]], Dict[str, ValueType]]): загружаемый model state_dict. Если ключом model_state_dict является nn.Module, этот ключ — подмодуль model, а значение должно быть state_dict подмодуля. При загрузке state_dict к нему будет добавлен префикс подмодуля.
  • optim_state_dict (OptimizerStateType) – OptimizerStateType: загружаемый state_dict оптимизатора.
  • options (StateDictOptions) – параметры, управляющие загрузкой model state_dict и optimizer state_dict. Подробности см. в StateDictOptions.
Возвращает:
  • missing_keys — список строк с отсутствующими ключами model state_dict.
  • unexpected_keys — список строк с неожиданными ключами model state_dict.
Тип возвращаемого значения:

NamedTuple с полями missing_keys и unexpected_keys

torch.distributed.checkpoint.state_dict.set_model_state_dict(model, model_state_dict, *, options=None) [исходный код]

Загружает state_dict модели.

Метод, обратный get_model_state_dict, который устанавливает state_dict для модели. Подробные сведения об использовании см. в set_state_dict.

Параметры:
  • model (nn.Module) – nn.Module модели.
  • model_state_dict (Dict[str, ValueType]) – (Dict[str, ValueType]): загружаемый model state_dict. Если ключом model_state_dict является nn.Module, этот ключ — подмодуль model, а значение должно быть state_dict подмодуля. При загрузке state_dict к нему будет добавлен префикс подмодуля.
  • options (StateDictOptions) – параметры, управляющие загрузкой model state_dict и optimizer state_dict. Подробности см. в StateDictOptions.
Возвращает:
  • missing_keys — список строк с отсутствующими ключами
  • unexpected_keys — список строк с неожиданными ключами
Тип возвращаемого значения:

NamedTuple с полями missing_keys и unexpected_keys

torch.distributed.checkpoint.state_dict.set_optimizer_state_dict(model, optimizers, optim_state_dict, *, options=None) [исходный код]

Загружает state_dict оптимизаторов.

Метод, обратный get_optimizer_state_dict, который устанавливает state_dict для оптимизаторов. Подробные сведения об использовании см. в set_state_dict.

WARN: set_optimizer_state_dict can only be called before backward() or after

step() вызывается для оптимизаторов. В противном случае состояния оптимизатора не будут корректно инициализированы.

Параметры:
  • model (nn.Module) – nn.Module модели.
  • optimizers (Union[Optimizer, Iterable[Optimizer]]) – оптимизаторы, используемые для оптимизации model.
  • optim_state_dict (OptimizerStateType) – OptimizerStateType: загружаемый state_dict оптимизатора.
  • options (StateDictOptions) – параметры, управляющие загрузкой model state_dict и optimizer state_dict. Подробности см. в StateDictOptions.
Возвращает:

None

Тип возвращаемого значения:

None

class torch.distributed.checkpoint.state_dict.StateDictOptions(full_state_dict=False, cpu_offload=False, ignore_frozen_params=False, keep_submodule_prefixes=True, strict=True, broadcast_from_rank0=False, flatten_optimizer_state_dict=False, dsd_fqn_modifiers='_fqn_modifiers') [исходный код]

Этот класс данных определяет поведение get_state_dict/set_state_dict.

  • full_state_dict: если установлено значение True, все тензоры в возвращаемом state_dict будут собраны. В возвращаемом state_dict не будет ShardedTensor и DTensor.
  • cpu_offload: выгрузить все тензоры в CPU. Чтобы избежать нехватки памяти на CPU, если full_state_dict также имеет значение True, state_dict получит только ранг 0, а все остальные ранги получат пустой state_dict.
  • ignore_frozen_params: если значение равно True, возвращаемый state_dict не будет содержать замороженных параметров — requires_grad имеет значение False. Значение по умолчанию — False.
  • keep_submodule_prefixes (устарело): если submodules не равно None, этот параметр указывает, нужно ли сохранять префиксы подмодулей в ключах state_dict. Например, если подмодуль — module.pretrain, а полный FQN параметра — pretrain.layer1.weight параметра. Если этот параметр равен True, ключ параметра в возвращаемом state_dict будет pretrain.layer1.weight. Если параметр равен False, ключ будет layer1.weight. Обратите внимание: если keep_submodule_prefixes имеет значение False, возможны конфликты FQN, поэтому в submodules должен быть только один подмодуль.
  • strict: параметр strict, используемый, когда set_state_dict вызывает model.load_state_dict().
  • broadcast_from_rank0: when the option is True, rank0 should receive a

    полный state_dict; тензоры в state_dict/optim_state_dict будут передаваться другим рангам по одному. Другие ранги получат тензоры и разделят их в соответствии с локальными фрагментами модели и оптимизатора. При использовании этого параметра необходимо установить full_state_dict в True. В настоящее время этот параметр поддерживает только DTensor, но не устаревший ShardedTensor.

Для пользователей, привыкших использовать модели и делиться ими в формате torch.save, предусмотрены следующие методы, предоставляющие автономные средства преобразования между форматами.

torch.distributed.checkpoint.format_utils.dcp_to_torch_save(dcp_checkpoint_dir, torch_save_path) [исходный код]

Преобразует каталог с контрольной точкой DCP в файл сохранения Torch.

Параметры:
  • dcp_checkpoint_dir (str | PathLike) – каталог с контрольной точкой DCP.
  • torch_save_path (str | PathLike) – имя файла для сохранения преобразованного файла Torch.

Предупреждение

Во избежание нехватки памяти рекомендуется запускать эту функцию только на одном ранге.

torch.distributed.checkpoint.format_utils.torch_save_to_dcp(torch_save_path, dcp_checkpoint_dir) [исходный код]

Преобразует файл сохранения Torch в контрольную точку DCP.

Параметры:
  • torch_save_path (str | PathLike) – имя файла сохранения Torch.
  • dcp_checkpoint_dir (str | PathLike) – каталог для сохранения контрольной точки DCP.

Предупреждение

Во избежание нехватки памяти рекомендуется запускать эту функцию только на одном ранге.

Следующие классы также можно использовать для загрузки и переразбиения моделей в оперативном режиме из формата torch.save.

class torch.distributed.checkpoint.format_utils.BroadcastingTorchSaveReader(checkpoint_id=None, coordinator_rank=0) [исходный код]

StorageReader для чтения файла Torch Save. Этот считыватель прочитает всю контрольную точку на ранге координатора, а затем выполнит широковещательную рассылку и разделит каждый тензор между всеми рангами.

. Примечание. Предназначен для использования с DynamicMetaLoadPlanner

Предупреждение

Текущая реализация поддерживает загрузку только тензоров.

>>> sd = {"mode": model}
>>> dcp.load(
>>>    sd,
>>>    storage_reader=BroadcastingTorchSaveReader(),
>>>    planner=DynamicMetaLoadPlanner(),
>>>    checkpoint_id="path_to_model.pt"
>>> )
prepare_global_plan(global_plan) [исходный код]

Реализация метода StorageReader

Тип возвращаемого значения:

list[LoadPlan]

prepare_local_plan(plan) [исходный код]

Реализация метода StorageReader

Тип возвращаемого значения:

LoadPlan

read_data(plan, planner) [исходный код]

Читает данные Torch Save на ранге координатора, а затем выполняет их широковещательную рассылку. Это влечёт за собой затраты на обмен данными, но позволяет не загружать всю контрольную точку на каждом ранге и, как ожидается, помогает избежать проблем с нехваткой памяти (OOM).

Тип возвращаемого значения:

Future[None]

read_metadata() [исходный код]

Расширяет StorageReader по умолчанию, чтобы обеспечить создание файла метаданных

Тип возвращаемого значения:

Metadata

reset(checkpoint_id=None) [исходный код]

Реализация метода StorageReader

set_up_storage_reader(metadata, is_coordinator) [исходный код]

Реализация метода StorageReader

classmethod validate_checkpoint_id(checkpoint_id) [исходный код]

Реализация метода StorageReader

Тип возвращаемого значения:

bool

class torch.distributed.checkpoint.format_utils.DynamicMetaLoadPlanner(flatten_state_dict=True, flatten_sharded_tensors=True, allow_partial_load=False) [исходный код]

Расширение DefaultLoadPlanner, которое создаёт новый объект Metadata на основе переданного state dict, избавляя от необходимости читать метаданные с диска. Это полезно при чтении форматов без файла метаданных, например файлов Torch Save.

. Примечание. Предназначен для использования с BroadcastingTorchSaveReader

Предупреждение

Текущая реализация поддерживает загрузку только тензоров.

>>> sd = {"mode": model}
>>> dcp.load(
>>>    sd,
>>>    storage_reader=BroadcastingTorchSaveReader(),
>>>    planner=DynamicMetaLoadPlanner(),
>>>    checkpoint_id="path_to_model.pt"
>>> )
set_up_planner(state_dict, metadata=None, is_coordinator=False) [исходный код]

Настройка планировщика: стандартное поведение расширяется созданием объекта Metadata из state dict

Для повышения наблюдаемости в производственных средах предоставлены следующие экспериментальные интерфейсы:

© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/distributed.checkpoint.html

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API