Распределённый чекпоинт — torch.distributed.checkpoint
Поддержка распределённого чекпоинта (DCP) позволяет загружать и сохранять модели с нескольких рангов параллельно. Он обрабатывает перераспределение ресурсов во время загрузки, что позволяет сохранять модель в одной топологии кластера и загружать её в другую.
DCP отличается от torch.save и torch.load несколькими важными аспектами:
- Он генерирует несколько файлов на чекпоинт, по одному на ранг как минимум.
- Он работает на месте, то есть модель должна сначала выделить память для своих данных, и DCP использует эту память вместо выделения новой.
Точки входа для загрузки и сохранения чекпоинта следующие:
-
torch.distributed.checkpoint.load_state_dict(state_dict, storage_reader, process_group=None, coordinator_rank=0, no_dist=False, planner=None)[source] -
Загружает распределённый
state_dictв стиле SPMD.Каждый ранг попытается прочитать наименьшее количество данных, необходимое для удовлетворения запрошенного
state_dict. При загрузкеShardedTensorэкземпляров каждый ранг считывает только данные для своих локальных фрагментов.Предупреждение
Все тензоры в
state_dictдолжны быть выделены на их целевом устройстве до вызова этой функции.Все данные, не являющиеся тензорами, загружаются с помощью
torch.load()и изменяются на месте в state_dict.Предупреждение
Пользователи должны вызвать
load_state_dictна корневом модуле, чтобы гарантировать, что обработка после загрузки и данные, не являющиеся тензорами, правильно передаются.- Параметры
-
- state_dict (Dict[str, Any]) – State_dict для загрузки. Обратите внимание, что этот state_dict будет обновлён на месте.
- storage_reader (StorageReader) – StorageReader, используемый для загрузки данных.
- process_group (ProcessGroup) – ProcessGroup, используемый для межранговой синхронизации.
- coordinator_rank (int) – Ранг, используемый для координации чекпоинта. По умолчанию используется ранг 0.
-
no_dist (bool) – Если
True, распределённый чекпоинт не будет сохраняться в стиле SPMD. (По умолчанию: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 использует коллективы для координации чтения между рангами. Для process group на основе NCCL внутренние представления тензоров объектов должны быть перемещены на устройство GPU перед началом коммуникации. В этом случае используемое устройство задаётся параметром
torch.cuda.current_device(), и пользователь несёт ответственность за его настройку таким образом, чтобы каждый ранг имел свой собственный GPU, с помощьюtorch.cuda.set_device().
-
torch.distributed.checkpoint.save_state_dict(state_dict, storage_writer, process_group=None, coordinator_rank=0, no_dist=False, planner=None)[source] -
Сохраняет распределённую модель в стиле SPMD.
Эта функция отличается от
torch.save()тем, что обрабатываетShardedTensorпутём сохранения каждым рангом только своих локальных фрагментов.Предупреждение
Нет гарантий обратной совместимости между версиями PyTorch для сохранённых state_dict.
Предупреждение
Если используется аргумент
process_group, убедитесь, что только его ранги вызываютsave_state_dict, и все данные в state_dict относятся к нему.Примечание
При сохранении чекпоинта для
ShardingStrategy.HYBRID_SHARDFSDP только один из групп фрагментов должен вызыватьsave_state_dict, и соответствующий process group должен быть передан.Примечание
Эта функция может быть использована для сохранения state_dict без инициализации process group, передавая
no_dist=True.- Параметры
-
- state_dict (Dict[str, Any]) – State_dict для сохранения.
- storage_writer (StorageWriter) – Экземпляр StorageWrite для выполнения записи.
- process_group (ProcessGroup) – ProcessGroup, используемый для межранговой синхронизации.
- coordinator_rank (int) – Ранг, используемый для координации чекпоинта. По умолчанию используется ранг 0.
-
no_dist (bool) – Если
True, распределённый чекпоинт не будет сохраняться в стиле SPMD. (По умолчанию:False)
- Возвращает
-
Объект метаданных для сохранённого чекпоинта.
- Тип возвращаемого значения
-
Metadata
Пример
>>> my_model = MyModule()
>>> model_state_dict = my_model.state_dict()
>>> fs_storage_writer = torch.distributed.checkpoint.FileSystemWriter("/checkpoint/1") >>> torch.distributed.checkpoint.save_state_dict( >>> state_dict=model_state_dict, >>> storage_writer=fs_storage_writer, >>> )Примечание
save_state_dict использует коллективы для координации записей между рангами. Для process group на основе NCCL внутренние представления тензоров объектов должны быть перемещены на устройство GPU перед началом коммуникации. В этом случае используемое устройство задаётся параметром
torch.cuda.current_device(), и пользователь несёт ответственность за его настройку таким образом, чтобы каждый ранг имел свой собственный GPU, с помощьюtorch.cuda.set_device().
Этот пример демонстрирует использование распределённого чекпоинта PyTorch для сохранения модели FSDP.
Следующие типы определяют интерфейс ввода-вывода, используемый во время чекпоинта:
-
class torch.distributed.checkpoint.StorageReader[source] -
Интерфейс, используемый
load_state_dictдля чтения из хранилища.Один экземпляр StorageReader действует как координатор и последователь в распределенном контрольной точке. В рамках инициализации каждому экземпляру сообщается его роль.
Подкласс должен ожидать следующую последовательность вызовов от
load_state_dict:- (все ранги) read_metadata()
- (все ранги) set_up_storage_reader()
- (все ранги) prepare_local_plan()
- (координатор) prepare_global_plan()
- (все ранги) read_data()
-
abstract prepare_global_plan(plans)[source] -
Выполнить централизованное планирование загрузки из хранилища.
Этот метод вызывается только на экземпляре координатора.
Хотя этот метод может создать совершенно другой план, предпочтительным способом является хранение данных, специфичных для хранилища, в LoadPlan::storage_data.
-
abstract prepare_local_plan(plan)[source] -
Выполнить локальное планирование, специфичное для хранилища.
Хотя этот метод может создать совершенно другой план, рекомендуемый способ - хранить данные, специфичные для хранилища, в LoadPlan::storage_data.
-
abstract read_data(plan, planner)[source] -
Читает все элементы из
plan, используяplannerдля разрешения данных.Подкласс должен вызвать
LoadPlanner::load_bytesдля десериализации объекта BytesIO в нужное место.Подкласс должен вызвать
LoadPlanner::resolve_tensorдля получения доступа к тензорам, в которые должны быть загружены данные.Ответственность StorageLayer заключается в правильном планировании любых необходимых копий между устройствами.
- Parameters
-
- plan (LoadPlan) – Локальный план для выполнения
- planner (LoadPlanner) – Объект планировщика для использования при разрешении элементов.
- Returns
-
Будущее, которое завершается, когда все чтения завершены.
- Return type
-
Future[None]
-
abstract read_metadata()[source] -
Читает метаданные контрольной точки.
- Returns
-
Объект метаданных, связанный с загружаемой контрольной точкой.
- Return type
-
Metadata
-
abstract set_up_storage_reader(metadata, is_coordinator)[source] -
Инициализировать этот экземпляр.
- Parameters
-
- metadata (Metadata) – Схема метаданных для использования.
- is_coordinator (bool) – Является ли этот экземпляр ответственным за координацию контрольной точки.
-
class torch.distributed.checkpoint.StorageWriter[source] -
Интерфейс, используемый
save_state_dictдля записи в хранилище.Один экземпляр StorageWriter выполняет функции как координатора, так и последователя при распределённом контрольно-восстановительном пункте. В процессе инициализации каждому экземпляру сообщается его роль.
Подкласс должен ожидать следующую последовательность вызовов.
- (все ранги) set_up_storage_writer()
- (все ранги) prepare_local_plan()
- (координатор) prepare_global_plan()
- (все ранги) write_data()
- (координатор) finish()
-
abstract finish(metadata, results)[source] -
Записывает метаданные и отмечает текущий контрольно-восстановительный пункт как успешный.
Фактический формат/схема, используемая для сериализации
metadata— деталь реализации. Единственное требование — это возможность восстановления в ту же графу объектов.
-
abstract prepare_global_plan(plans)[source] -
Выполняет централизованное планирование хранения.
Этот метод вызывается только на экземпляре координатора.
Хотя этот метод может генерировать совершенно другой план, предпочтительным способом является хранение данных, специфичных для хранения, в SavePlan::storage_data.
-
abstract prepare_local_plan(plan)[source] -
Выполняет локальное планирование, специфичное для хранения.
Хотя этот метод может генерировать совершенно другой план, рекомендуемым способом является хранение данных, специфичных для хранения, в SavePlan::storage_data.
-
abstract set_up_storage_writer(is_coordinator)[source] -
Инициализирует этот экземпляр.
- Параметры
-
is_coordinator (bool) — является ли этот экземпляр ответственным за координацию контрольно-восстановительного пункта.
-
abstract write_data(plan, planner)[source] -
Записывает все элементы из
planс использованиемplannerдля разрешения данных.Подкласс должен вызывать
SavePlanner::resolve_dataдля каждого элемента из плана, чтобы получить доступ к базовому объекту для записи.Подклассы должны лениво вызывать
resolve_data, так как это может выделять память. В случае тензоров следует сделать следующие предположения:- Они могут находиться на любом устройстве, включая устройство, не соответствующее устройству
WriteItem::tensor_data - Они могут быть представлениями или не сплошными. Должен быть сохранен только проекции.
- Параметры
-
- plan (SavePlan) — план сохранения для выполнения.
- planner (SavePlanner) — объект планировщика, используемый для разрешения элементов до данных.
- Возвращает
-
Будущее, которое завершается списком WriteResult
- Тип возвращаемого значения
- Они могут находиться на любом устройстве, включая устройство, не соответствующее устройству
Следующие типы определяют интерфейс планировщика, используемый во время контрольно-восстановительного пункта:
-
class torch.distributed.checkpoint.LoadPlanner[source] -
Абстрактный класс, определяющий протокол, используемый load_state_dict для планирования процесса загрузки.
Объекты LoadPlanner являются объектами состояния, которые могут быть использованы для настройки всего процесса загрузки.
LoadPlanner действует как прокси-сервер доступа к state_dict, поэтому любые преобразования, выполненные с ним, будут видны всему процессу.
Подкласс LoadPlanner может ожидать следующую последовательность вызовов во время load_state_dict:
-
- set_up_planner — вызывается на всех рангах.
-
Сигнализирует о начале загрузки контрольной точки.
-
- create_local_plan — вызывается на всех рангах.
-
Обрабатывает state_dict и создаёт
LoadPlan, который будет отправлен для глобального планирования.
-
- create_global_plan — вызывается только на координирующем ранге.
-
Принимает LoadPlan со всех рангов и принимает любые глобальные решения.
-
- load_bytes — вызывается многократно на каждом ранге
-
Вызывается один раз для каждого не-тензорного значения в state_dict.
-
- resolve_tensor и commit_tensor — вызываются многократно на каждом ранге
-
Вызываются парами для каждого тензорного значения в state_dict.
Пользователям рекомендуется расширять DefaultLoadPlanner вместо этого интерфейса напрямую, так как большинство изменений могут быть выражены изменениями в единственном методе.
Существует два типичных шаблона расширения:
Перезапись state_dict. Это самый простой способ расширить процесс загрузки, так как он не требует понимания тонкостей работы LoadPlan. Нам нужно сохранить ссылку на исходный state_dict, так как загрузка происходит на месте, поэтому нам нужно иметь возможность выполнить её на месте.
>>> class RenamePlanner(DefaultLoadPlanner): >>> def set_up_planner(self, state_dict, metadata, is_coordinator): >>> self.original_state_dict = state_dict >>> super().set_up_planner(self, {"foo_" + k: v for k, v in state_dict.items()}, 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)Изменение 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)[source] -
Этот метод вызывается после того, как StorageReader завершил загрузку данных в
tensor.Предоставленный тензор — тот же, что и возвращаемый вызовом
resolve_tensor. Этот метод необходим только в том случае, если этому LoadPlanner необходимо выполнить пост-обработкуtensorперед копированием обратно в тот, что в state_dict.Содержимое тензора будет следовать своей модели синхронизации устройств.
-
abstract create_global_plan(global_plan)[source] -
Вычислить глобальный план загрузки и вернуть планы для каждого ранга.
. Примечание. Этот метод вызывается только на координирующем ранге
-
abstract create_local_plan()[source] -
Создать LoadPlan на основе state_dict и метаданных, предоставленных set_up_planner.
. Примечание. Этот метод вызывается на каждом ранге.
- Тип возвращаемого значения
-
abstract finish_plan(central_plan)[source] -
Принять план от координатора и вернуть окончательный LoadPlan.
- Тип возвращаемого значения
-
abstract load_bytes(read_item, value)[source] -
Загрузить элемент, описанный
read_item``and ``value.Этот метод должен изменять на месте базовый state_dict.
Содержимое
valueопределяется SavePlanner, используемым для создания загружаемой контрольной точки.
-
abstract resolve_tensor(read_item)[source] -
Вернуть тензор, описанный
read_item, для использования StorageReader для загрузкиread_item.Тензор должен быть связан с тем, что в базовом state_dict, так как StorageReader заменит его содержимое. Если по какой-либо причине это невозможно, планировщик может использовать метод
commit_tensorдля копирования данных обратно в тот, что в state_dict.- Тип возвращаемого значения
-
abstract set_up_planner(state_dict, metadata, is_coordinator)[source] -
Инициализировать этот экземпляр для загрузки данных в
state_dict.. Примечание. Этот метод вызывается на каждом ранге.
-
-
class torch.distributed.checkpoint.LoadPlan(items: List[torch.distributed.checkpoint.planner.ReadItem], storage_data: Any = None, planner_data: Any = None)[source]
-
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)[source]
-
class torch.distributed.checkpoint.SavePlanner[source] -
Абстрактный класс, определяющий протокол, используемый save_state_dict для планирования процесса сохранения.
SavePlanners — это объекты состояния, которые можно использовать для настройки всего процесса сохранения.
SavePlanner действует как прокси-сервер доступа к state_dict, поэтому любые преобразования, выполненные с ним, будут видны всему процессу.
Подкласс SavePlanner может ожидать следующую последовательность вызовов во время 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, is_coordinator): >>> # prefix all keys with `foo_`` >>> super().set_up_planner({"foo_" + k: v for k, v in state_dict.items()}, 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 islice >>> 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): >>> def chunk(it, size): >>> it = iter(it) >>> return list(iter(lambda: tuple(islice(it, size)), ())) >>> all_plans = [ >>> replace(plan, items=items) for plan, items in >>> zip(all_plans, chunk(all_plans[0].items, len(all_plans))) >>> ] >>> 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)[source] -
Вычислить глобальный план контрольной точки и вернуть локальный план каждого ранга.
Вызывается только на ранге координатора.
-
abstract create_local_plan()[source] -
Вычислить план сохранения для текущего ранга. Он будет агрегирован и передан в create_global_plan. Данные, специфичные для планировщика, могут передаваться через SavePlan::planner_data.
Вызывается на всех рангах.
- Return type
-
abstract finish_plan(new_plan)[source] -
Объединить план, созданный
create_local_plan, и результатcreate_global_plan.Вызывается на всех рангах.
- Return type
-
abstract resolve_data(write_item)[source] -
Ищет объект, связанный с
write_itemвstate_dict, и применяет любое преобразование (например, сериализацию) перед тем, как хранилище его использует.Вызывается на каждом ранге несколько раз, по крайней мере, один раз на каждый WriteItem в окончательном SavePlan.
Этот метод должен быть идемпотентным и потокобезопасным. Реализации StorageWriter могут вызывать его так часто, как им нужно.
Любое преобразование, выделяющее память, должно выполняться лениво при вызове этого метода, чтобы уменьшить пиковую потребность в памяти при создании контрольной точки.
При возвращении тензоров они могут находиться на любом устройстве или в любом формате, они также могут быть представлениями. Хранилище отвечает за то, как их сохранить.
- Return type
-
Объединение[Tensor, BytesIO]
-
abstract set_up_planner(state_dict, is_coordinator)[source] -
Инициализировать этот планировщик для сохранения
state_dict.Реализации должны сохранить эти значения, так как они не будут предоставлены позже в процессе сохранения.
Вызывается на всех рангах.
-
-
class torch.distributed.checkpoint.SavePlan(items: List[torch.distributed.checkpoint.planner.WriteItem], storage_data: Any = None, planner_data: Any = None)[source]
-
class torch.distributed.checkpoint.WriteItem(index: torch.distributed.checkpoint.metadata.MetadataIndex, type: torch.distributed.checkpoint.planner.WriteItemType, tensor_data: Union[torch.distributed.checkpoint.planner.TensorWriteData, NoneType] = None)[source]
Мы предоставляем хранилище на основе файловой системы:
-
class torch.distributed.checkpoint.FileSystemReader(path)[source]
-
class torch.distributed.checkpoint.FileSystemWriter(path, single_file_per_rank=True, sync_files=True, thread_count=1, per_thread_copy_ahead=10000000)[source] -
Базовая реализация StorageWriter с использованием файлового ввода-вывода.
Эта реализация делает следующие предположения и упрощения:
- Путь к контрольной точке — это пустая или несуществующая директория.
- Создание файла атомарно.
Контрольная точка состоит из одного файла на запрос записи плюс файл
.metadataс сериализованными метаданными.
Мы предоставляем стандартные реализации LoadPlanner и SavePlanner, которые могут обрабатывать все конструкции torch.distributed, такие как FSDP, DDP, ShardedTensor и DistributedTensor.
-
class torch.distributed.checkpoint.DefaultSavePlanner(flatten_state_dict=True, flatten_sharded_tensors=True, dedup_replicated_tensors=True)[source] -
-
lookup_object(index)[source] -
Это расширение от интерфейса планировщика, чтобы облегчить расширение стандартного планировщика
- Return type
-
transform_object(write_item, object)[source] -
Это расширение от интерфейса планировщика, чтобы облегчить расширение стандартного планировщика
-
-
class torch.distributed.checkpoint.DefaultLoadPlanner(flatten_state_dict=True, flatten_sharded_tensors=True)[source] -
Планировщик DefaultLoadPlanner, добавляющий несколько функций поверх LoadPlanner.
В частности, он добавляет следующее:
flatten_state_dict: Обработка state_dict со вложенными словарями flatten_sharded_tensors: Для FSDP в режиме 2D параллельности
-
lookup_tensor(index)[source] -
Это расширение интерфейса планировщика, чтобы упростить расширение стандартного планировщика.
- Тип возвращаемого значения
-
transform_tensor(read_item, tensor)[source] -
Это расширение интерфейса планировщика, чтобы упростить расширение стандартного планировщика.
-
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/distributed.checkpoint.html