Функция get_rank
Функция distributed.get_rank используется для получения ранга текущего процесса
в распределенной группе. Ранг - это целое число, которое уникально идентифицирует
каждый процесс в группе. Ранг принимает значения от 0 до world_size - 1,
где world_size - это общее количество процессов в группе.
Функция не принимает параметров и возвращает целое число.
Ранг процесса используется для управления поведением программы в зависимости от того,
какой процесс выполняет код, например, для вывода логов только с главного процесса
или для распределения данных между процессами.
Синтаксис
import torch.distributed as dist
rank = dist.get_rank()
Пример
Перед вызовом get_rank необходимо инициализировать группу процессов
с помощью init_process_group. Рассмотрим базовый пример запуска скрипта
с двумя процессами и получения их рангов:
import torch
import torch.distributed as dist
# Инициализация группы процессов для двух процессов
dist.init_process_group(backend='gloo', init_method='tcp://localhost:23456', world_size=2, rank=0)
rank = dist.get_rank()
print(f"Current process rank: {rank}")
dist.destroy_process_group()
Результат выполнения кода для процесса с рангом 0:
"Current process rank: 0"
Результат выполнения кода для процесса с рангом 1:
"Current process rank: 1"
Пример
Часто get_rank используется для выполнения действий только на главном процессе
(с рангом 0), например, для сохранения модели или вывода информации:
import torch
import torch.distributed as dist
dist.init_process_group(backend='gloo', init_method='tcp://localhost:23456', world_size=4, rank=0)
rank = dist.get_rank()
if rank == 0:
print("Saving model checkpoint from main process...")
# Здесь может быть код сохранения модели
else:
print(f"Process {rank} is waiting for main process...")
dist.destroy_process_group()
Результат выполнения кода для процесса с рангом 0:
"Saving model checkpoint from main process..."
Результат выполнения кода для процесса с рангом 2:
"Process 2 is waiting for main process..."
Пример
Ранг может использоваться для загрузки разных частей данных на разные процессы при распределенном обучении. Например, можно разбить датасет на части в соответствии с рангом процесса:
import torch
import torch.distributed as dist
dist.init_process_group(backend='gloo', init_method='tcp://localhost:23456', world_size=4, rank=1)
total_data = 100
world_size = dist.get_world_size()
rank = dist.get_rank()
# Вычисление индексов для текущего процесса
chunk_size = total_data // world_size
start_idx = rank * chunk_size
end_idx = start_idx + chunk_size
print(f"Process {rank} processes data from index {start_idx} to {end_idx-1}")
dist.destroy_process_group()
Результат выполнения кода для процесса с рангом 1:
"Process 1 processes data from index 25 to 49"
Смотрите также
-
функцию
get_world_size,
которая возвращает общее количество процессов в группе -
функцию
init_process_group,
которая инициализирует распределенную группу процессов -
функцию
destroy_process_group,
которая уничтожает группу процессов -
функцию
is_initialized,
которая проверяет, инициализирована ли группа процессов