Spec-Zone.ru › PyTorch 2

Распределённый чекпоинт — 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_SHARD FSDP только один из групп фрагментов должен вызывать 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:

  1. (все ранги) read_metadata()
  2. (все ранги) set_up_storage_reader()
  3. (все ранги) prepare_local_plan()
  4. (координатор) prepare_global_plan()
  5. (все ранги) read_data()
abstract prepare_global_plan(plans) [source]

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

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

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

Parameters

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

Returns

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

Return type

List[LoadPlan]

abstract prepare_local_plan(plan) [source]

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

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

Parameters

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

Returns

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

Return type

LoadPlan

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 выполняет функции как координатора, так и последователя при распределённом контрольно-восстановительном пункте. В процессе инициализации каждому экземпляру сообщается его роль.

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

  1. (все ранги) set_up_storage_writer()
  2. (все ранги) prepare_local_plan()
  3. (координатор) prepare_global_plan()
  4. (все ранги) write_data()
  5. (координатор) finish()
abstract finish(metadata, results) [source]

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

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

Параметры
  • metadata (Metadata) — метаданные для нового контрольно-восстановительного пункта
  • results (Список[Список[WriteResult]]) — список WriteResults со всех рангов.
Возвращает

None

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

None

abstract prepare_global_plan(plans) [source]

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

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

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

Параметры

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

Возвращает

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

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

Список[SavePlan]

abstract prepare_local_plan(plan) [source]

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

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

Параметры

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

Возвращает

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

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

SavePlan

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

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

Future[Список[WriteResult]]

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

class torch.distributed.checkpoint.LoadPlanner [source]

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

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

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

Подкласс LoadPlanner может ожидать следующую последовательность вызовов во время 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 — вызываются многократно на каждом ранге

    Вызываются парами для каждого тензорного значения в 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]

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

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

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

Список[LoadPlan]

abstract create_local_plan() [source]

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

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

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

LoadPlan

abstract finish_plan(central_plan) [source]

Принять план от координатора и вернуть окончательный LoadPlan.

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

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.

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

Tensor

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:

  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, 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]

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

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

Return type

Tuple[Список[SavePlan], Метаданные]

abstract create_local_plan() [source]

Вычислить план сохранения для текущего ранга. Он будет агрегирован и передан в create_global_plan. Данные, специфичные для планировщика, могут передаваться через SavePlan::planner_data.

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

Return type

SavePlan

abstract finish_plan(new_plan) [source]

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

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

Return type

SavePlan

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

Spec-Zone.ru

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