Распределённые контрольные точки — 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)
-
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]) – локальные фрагменты, которые необходимо загрузить.
-
fqn (str) – FQN элемента state_dict для передачи в
- Возвращает:
-
Список
ReadItem, которые удовлетворят требованиям всех входных фрагментов. - Тип возвращаемого значения:
-
torch.distributed.checkpoint.default_planner.create_default_global_load_plan(all_plans)[исходный код] -
Создать глобальный план загрузки, используемый DefaultLoadPlanner.
По умолчанию загрузка не требует глобальной координации, и в настоящее время эта функция не изменяет локальные планы.
-
torch.distributed.checkpoint.default_planner.create_default_global_save_plan(all_plans, rewrite_index_hints=True)[исходный код] -
Создать глобальный план и метаданные, используемые DefaultSavePlanner.
Метаданные формируются путём объединения метаданных всех
WriteItemиз предоставленных планов.Единственное изменение при глобальном планировании — обновление подсказок индекса во всех объектах
MetadataIndex, еслиrewrite_index_hintsравно True.
-
torch.distributed.checkpoint.default_planner.create_default_local_save_plan(state_dict, is_coordinator)[исходный код] -
Создать
SavePlan, используемый DefaultSavePlanner.На рангах, не являющихся координатором, эта функция игнорирует тензоры и объекты, не являющиеся тензорами, и создаёт операции записи только для объектов ShardedTensor.
На ранге координатора создаются операции записи для всех значений.
- Тип возвращаемого значения:
-
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_SHARDFSDP функцию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. - Тип возвращаемого значения:
Пример
>>> 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. Несовпадение ключей может привести к зависанию или ошибкам. Если вы не уверены, можно использовать APIutils._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) выглядит следующим образом:-
- AsyncStager.stage_data(state_dict):
-
Этот вызов предоставляет AsyncStager возможность подготовить state_dict. Цель подготовки в данном контексте — создать представление state_dict, безопасное для обучения: любые изменения данных модуля после завершения подготовки не должны отражаться в state_dict, возвращаемом этим методом. Например, в стандартном случае на CPU создаётся копия всего state_dict, что позволяет продолжать обучение, не рискуя изменить данные, которые сериализуются.
-
- Для state_dict, возвращённого stage, параллельно вызывается dcp.save. Этот вызов отвечает
-
за сериализацию state_dict и его запись в хранилище.
-
- Если 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.
-
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.
-
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.
-
synchronize_staging()[исходный код] -
Функция ничего не делает, поскольку подготовка выполняется блокирующим образом.
-
Помимо описанных выше точек входа, объекты Stateful, описание которых приведено ниже, позволяют дополнительно настраивать сохранение и загрузку.
-
class torch.distributed.checkpoint.stateful.Stateful(*args, **kwargs)[исходный код] -
Протокол Stateful для объектов, которые можно сохранять в контрольную точку и восстанавливать из неё.
-
load_state_dict(state_dict)[исходный код] -
Восстанавливает состояние объекта из предоставленного state_dict.
-
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:- (все ранги) устанавливает checkpoint_id, если пользователи передали допустимый checkpoint_id.
- (все ранги) вызывает read_metadata()
- (все ранги) вызывает set_up_storage_reader()
- (все ранги) вызывает prepare_local_plan()
- (координатор) вызывает prepare_global_plan()
- (все ранги) вызывает read_data()
-
abstract prepare_global_plan(plans)[исходный код] -
Выполняет централизованное планирование загрузки из хранилища.
Этот метод вызывается только для экземпляра координатора.
Хотя этот метод может сформировать совершенно иной план, предпочтительный способ — сохранять специфичные для хранилища данные в LoadPlan::storage_data.
-
abstract prepare_local_plan(plan)[исходный код] -
Выполняет локальное планирование, специфичное для хранилища.
Хотя этот метод может сформировать совершенно иной план, рекомендуемый способ — сохранять специфичные для хранилища данные в LoadPlan::storage_data.
-
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 хранилищем. Это позволяет автоматически выбирать хранилище.
- Тип возвращаемого значения:
-
class torch.distributed.checkpoint.StorageWriter[исходный код] -
Интерфейс, используемый
save_state_dictдля записи в хранилище.Один экземпляр StorageWriter выступает как координатор и как ведомый в распределенной контрольной точке. При инициализации каждому экземпляру назначается роль.
Подкласс должен ожидать следующую последовательность вызовов.
- (все ранги) задают checkpoint_id, если пользователи передали допустимый checkpoint_id.
- (все ранги) вызывают set_up_storage_writer()
- (все ранги) вызывают prepare_local_plan()
- (координатор) вызывает prepare_global_plan()
- (все ранги) вызывают write_data()
- (координатор) вызывает finish()
-
abstract finish(metadata, results)[исходный код] -
Записывает метаданные и помечает текущую контрольную точку как успешно созданную.
Фактический формат/схема, используемые для сериализации
metadata, являются деталями реализации. Единственное требование — возможность восстановить тот же граф объектов.
-
abstract prepare_global_plan(plans)[исходный код] -
Выполняет централизованное планирование хранения.
Этот метод вызывается только для экземпляра координатора.
Хотя этот метод может создавать совершенно другой план, предпочтительно сохранять специфичные для хранилища данные в SavePlan::storage_data.
-
abstract prepare_local_plan(plan)[исходный код] -
Выполняет локальное планирование с учетом особенностей хранилища.
Хотя этот метод может создавать совершенно другой план, рекомендуется сохранять специфичные для хранилища данные в SavePlan::storage_data.
-
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 хранилищем. Это позволяет автоматически выбирать хранилище.
- Тип возвращаемого значения:
-
abstract write_data(plan, planner)[исходный код] -
Записывает все элементы из
plan, используяplannerдля получения данных.Подкласс должен вызывать
SavePlanner::resolve_dataдля каждого элемента плана, чтобы получить доступ к базовому объекту для записи.Подклассам следует вызывать
resolve_dataпо мере необходимости, поскольку это может выделять память. Для тензоров следует учитывать следующее:- Они могут находиться на любом устройстве, в том числе не совпадающем с устройством в
WriteItem::tensor_data - Они могут быть представлениями или неконт contiguousными. Сохранять нужно только проекцию.
- Параметры:
-
- plan (SavePlan) – план сохранения для выполнения.
- planner (SavePlanner) – объект планировщика, используемый для получения данных из элементов.
- Возвращает:
-
объект future, который после завершения возвращает список 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:
-
- set_up_planner — вызывается на всех рангах.
-
Сигнализирует о начале загрузки контрольной точки.
-
- create_local_plan — вызывается на всех рангах.
-
Обрабатывает state_dict и создает
LoadPlan, который будет передан для глобального планирования.
-
- create_global_plan — вызывается только на ранге координатора.
-
Получает LoadPlan со всех рангов и принимает глобальные решения.
-
- load_bytes — вызывается несколько раз на каждом ранге
-
Вызывается один раз для каждого значения в state_dict, не являющегося тензором.
-
- 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)[исходный код] -
Вычисляет глобальный план загрузки и возвращает планы для каждого ранга.
. Примечание: вызывается только на ранге координатора
-
abstract create_local_plan()[исходный код] -
Создает LoadPlan на основе state_dict и метаданных, переданных в set_up_planner.
. Примечание: вызывается на каждом ранге.
- Тип возвращаемого значения:
-
abstract finish_plan(central_plan)[исходный код] -
Принимает план от координатора и возвращает итоговый 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.- Тип возвращаемого значения:
-
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:
-
- set_up_planner — вызывается на всех рангах.
-
Сигнализирует о начале сохранения контрольной точки.
-
- create_local_plan — вызывается на всех рангах.
-
Обрабатывает state_dict и создает
SavePlan, который будет передан для глобального планирования.
-
- create_global_plan — вызывается только на ранге координатора.
-
Получает SavePlan со всех рангов и принимает глобальные решения.
-
- finish_plan — вызывается на всех рангах.
-
Позволяет каждому рангу скорректировать план с учетом решений глобального планирования.
-
- 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)[исходный код] -
Вычисляет глобальный план контрольной точки и возвращает локальный план каждого ранга.
Вызывается только на ранге координатора.
-
abstract create_local_plan()[исходный код] -
Вычисляет план сохранения для текущего ранга.
Этот план будет объединен и передан в create_global_plan. Специфичные для планировщика данные можно передать через SavePlan::planner_data.
Вызывается на всех рангах.
- Тип возвращаемого значения:
-
abstract finish_plan(new_plan)[исходный код] -
Объединяет план, созданный методом
create_local_plan, с результатомcreate_global_plan.Вызывается на всех рангах.
- Тип возвращаемого значения:
-
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с сериализованными метаданными.
Мы также предоставляем другие уровни хранения, в том числе для работы с 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)[исходный код] -
Расширение интерфейса планировщика, упрощающее расширение планировщика по умолчанию.
- Тип возвращаемого значения:
-
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)[исходный код] -
Расширение интерфейса планировщика, упрощающее расширение планировщика по умолчанию.
- Тип возвращаемого значения:
-
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. - Тип возвращаемого значения:
-
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. - Тип возвращаемого значения:
-
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.
- Параметры:
Предупреждение
Во избежание нехватки памяти рекомендуется запускать эту функцию только на одном ранге.
-
torch.distributed.checkpoint.format_utils.torch_save_to_dcp(torch_save_path, dcp_checkpoint_dir)[исходный код] -
Преобразует файл сохранения Torch в контрольную точку 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
-
prepare_local_plan(plan)[исходный код] -
Реализация метода StorageReader
- Тип возвращаемого значения:
-
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
- Тип возвращаемого значения:
-
-
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