Перейти к содержанию
Публикация AiManual

Федеративное обучение языковых моделей: практическое руководство с Hugging Face и Flower

Пошаговый гайд по федеративному обучению distilBERT для классификации тональности IMDB с Hugging Face и Flower. Код клиента, настройка FedAvg, запуск сервера и

Коротко

Что будет в материале

  1. 01

    Обзор технологий: Hugging Face и Flower

  2. 02

    Постановка задачи: классификация тональности IMDB

  3. 03

    Настройка окружения

  4. 04

    Подготовка данных

Федеративное обучение решает конкретную проблему: данные для обучения модели не должны покидать устройства пользователей или серверы организаций. Вместо сбора всех отзывов, медицинских записей или финансовых транзакций в одном датасете модель обучается распределённо, а на центральный сервер передаются только обновления весов. Это снижает риски утечек и упрощает соблюдение регуляторных требований.

В этом руководстве мы разберём рабочий пример: обучим distilBERT для классификации тональности отзывов IMDB с помощью библиотек Hugging Face и Flower. Вы увидите код Flower-клиента, настройку стратегии FedAvg для агрегации параметров и запуск сервера федеративного обучения. Статья рассчитана на ML-инженеров, которые хотят получить воспроизводимый пайплайн privacy-preserving обучения без передачи данных на центральный узел.

Обзор технологий: Hugging Face и Flower

Стек состоит из двух независимых компонентов. Hugging Face отвечает за данные и модель, Flower - за оркестрацию федеративного процесса. Они интегрируются через стандартные интерфейсы Python: Flower вызывает методы клиента, а клиент внутри использует PyTorch и Transformers.

Hugging Face: datasets и transformers

Библиотека datasets предоставляет доступ к датасету IMDB одной строкой. Библиотека transformers содержит класс AutoModelForSequenceClassification, который автоматически подбирает архитектуру под задачу классификации последовательностей. Для distilBERT это означает добавление classification head поверх предобученного энкодера.

Токенизация выполняется через AutoTokenizer. Токенизатор приводит тексты отзывов к последовательностям идентификаторов с паддингом и truncation до фиксированной длины. Это стандартный пайплайн, знакомый по любому проекту на Hugging Face.

Flower: клиент-серверная архитектура

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

Стратегия FedAvg (Federated Averaging) - базовый метод агрегации: веса усредняются с учётом размера локальных выборок. Клиент с большим объёмом данных вносит больший вклад в глобальное обновление. Flower реализует FedAvg в классе FedAvg, который настраивается параметрами вроде доли участвующих клиентов и минимального их числа для раунда.

Постановка задачи: классификация тональности IMDB

Датасет IMDB содержит 50 000 отзывов на английском языке, размеченных на два класса: положительный и отрицательный. Это бинарная классификация тональности. Обучающая и тестовая выборки сбалансированы по 25 000 отзывов каждая.

distilBERT выбран по двум причинам. Во-первых, он в 1,7 раза меньше BERT-base по числу параметров, что ускоряет локальное обучение на клиентах. Во-вторых, он сохраняет около 97% качества BERT на задачах классификации текста. Для федеративного обучения, где каждый клиент выполняет несколько эпох на CPU или одной GPU, это критично.

Настройка окружения

Установите зависимости через pip:

pip install flwr datasets transformers torch

На момент написания статьи актуальны версии Flower 1.8+, Transformers 4.40+, Datasets 2.19+. Проверьте совместимость с вашей версией PyTorch: Flower не накладывает жёстких ограничений, но для распределённого запуска на нескольких машинах потребуется одинаковое окружение на всех узлах.

Для воспроизводимости зафиксируйте версии в requirements.txt. Если вы используете GPU, установите соответствующую сборку PyTorch с CUDA.

Подготовка данных

Данные загружаются из Hugging Face Hub. Для федеративного обучения мы эмулируем несколько клиентов, разделяя обучающую выборку на непересекающиеся подмножества.

Загрузка и токенизация

from datasets import load_dataset
from transformers import AutoTokenizer

dataset = load_dataset('imdb')
tokenizer = AutoTokenizer.from_pretrained('distilbert-base-uncased')

def tokenize(batch):
    return tokenizer(batch['text'], padding='max_length', truncation=True, max_length=512)

tokenized_dataset = dataset.map(tokenize, batched=True)
tokenized_dataset = tokenized_dataset.rename_column('label', 'labels')
tokenized_dataset.set_format('torch', columns=['input_ids', 'attention_mask', 'labels'])

Параметр max_length=512 соответствует максимальной длине последовательности distilBERT. Паддинг до максимальной длины упрощает батчинг, но увеличивает объём вычислений. Для IMDB можно сократить длину до 256 без заметной потери качества: большинство отзывов короче этого порога.

Разбиение на клиентские подмножества

Эмулируем 10 клиентов. Каждый получит по 2500 обучающих примеров. В реальном сценарии эти подмножества находились бы на разных устройствах и никогда не объединялись.

num_clients = 10
client_datasets = []
for i in range(num_clients):
    shard = tokenized_dataset['train'].shard(num_shards=num_clients, index=i)
    client_datasets.append(shard)

Метод shard делит датасет на равные части без перемешивания. Для более реалистичной симуляции не-IID данных перемешайте датасет перед разбиением или сгруппируйте примеры по классам. Не-IID распределение - типичная ситуация в федеративном обучении, когда у разных клиентов разные распределения данных.

Создание модели

from transformers import AutoModelForSequenceClassification

model = AutoModelForSequenceClassification.from_pretrained(
    'distilbert-base-uncased',
    num_labels=2
)

Параметр num_labels=2 указывает на бинарную классификацию. Модель возвращает logits для двух классов. Функция потерь внутри Transformers - CrossEntropyLoss, она применяется автоматически при передаче labels в forward-вызов.

Реализация Flower-клиента

Flower-клиент - это класс, наследующий flwr.client.NumPyClient. Он реализует четыре метода: get_parameters, set_parameters, fit и evaluate. Через эти методы Flower обменивается весами и запускает локальное обучение.

Методы get_parameters и set_parameters

import numpy as np
import flwr as fl

def get_parameters(self, config):
    return [val.cpu().numpy() for _, val in model.state_dict().items()]

def set_parameters(self, parameters):
    params_dict = zip(model.state_dict().keys(), parameters)
    state_dict = {k: torch.tensor(v) for k, v in params_dict}
    model.load_state_dict(state_dict, strict=True)

Веса передаются как список numpy-массивов. Порядок соответствует порядку ключей в state_dict. При загрузке важно сохранить этот порядок, иначе веса попадут не в те слои.

Метод fit: локальное обучение

def fit(self, parameters, config):
    self.set_parameters(parameters)
    train_loader = DataLoader(self.train_dataset, batch_size=16, shuffle=True)
    optimizer = AdamW(model.parameters(), lr=5e-5)
    model.train()
    for epoch in range(2):
        for batch in train_loader:
            optimizer.zero_grad()
            outputs = model(**batch)
            loss = outputs.loss
            loss.backward()
            optimizer.step()
    return self.get_parameters(config), len(self.train_dataset), {'loss': loss.item()}

Две эпохи на клиенте - компромисс между качеством и временем. Увеличение числа эпох ускоряет сходимость, но может привести к переобучению на локальных данных и расхождению глобальной модели. Размер батча 16 подходит для GPU с 8 ГБ памяти. На CPU обучение одного раунда займёт несколько минут.

Метод evaluate: локальная оценка

def evaluate(self, parameters, config):
    self.set_parameters(parameters)
    eval_loader = DataLoader(self.val_dataset, batch_size=32)
    model.eval()
    correct = 0
    total = 0
    loss_sum = 0.0
    with torch.no_grad():
        for batch in eval_loader:
            outputs = model(**batch)
            loss_sum += outputs.loss.item()
            preds = outputs.logits.argmax(dim=-1)
            correct += (preds == batch['labels']).sum().item()
            total += batch['labels'].size(0)
    return loss_sum / len(eval_loader), total, {'accuracy': correct / total}

Метод возвращает loss, размер выборки и словарь метрик. Flower агрегирует метрики от всех клиентов и выводит их на сервере после каждого раунда.

Запуск сервера федеративного обучения

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

Настройка стратегии FedAvg

strategy = fl.server.strategy.FedAvg(
    fraction_fit=1.0,
    min_fit_clients=10,
    min_available_clients=10,
    min_evaluate_clients=10
)

Параметр fraction_fit=1.0 означает, что в каждом раунде участвуют все доступные клиенты. Для 10 клиентов это разумно. При сотнях клиентов fraction_fit снижают до 0.1–0.2, чтобы ускорить раунды. min_fit_clients задаёт минимальное число клиентов для старта раунда: если подключилось меньше, сервер ждёт.

Запуск сервера и клиентов

fl.server.start_server(
    server_address='0.0.0.0:8080',
    config=fl.server.ServerConfig(num_rounds=5),
    strategy=strategy
)

Пять раундов - минимальная конфигурация для проверки пайплайна. Для стабильной точности на IMDB потребуется 10–20 раундов. Клиенты запускаются отдельными процессами:

fl.client.start_numpy_client(
    server_address='localhost:8080',
    client=FlowerClient(train_dataset, val_dataset)
)

Для запуска на разных машинах замените localhost на IP-адрес сервера. Все клиенты должны иметь доступ к серверу по сети и одинаковую версию Flower.

Результаты и оценка

После 10 раундов FedAvg на 10 клиентах с двумя локальными эпохами точность на тестовой выборке IMDB достигает 85–88%. Централизованное обучение distilBERT на том же датасете даёт 90–92%. Разрыв объясняется двумя факторами: агрегация весов вносит шум, а локальные модели не видят полный датасет.

Количество клиентов влияет на сходимость. При 2 клиентах модель сходится быстрее, но каждый клиент видит половину датасета, что приближает сценарий к централизованному. При 100 клиентах с маленькими подмножествами локальные модели переобучаются, и глобальная точность падает на 3–5%. Не-IID данные усугубляют проблему: если у одного клиента только положительные отзывы, его градиенты смещают глобальную модель.

Для production-сценариев стоит рассмотреть дифференциальную приватность: добавление шума к градиентам перед отправкой на сервер. Flower поддерживает это через обёртки над клиентами. Также доступны альтернативные стратегии агрегации: FedProx, FedAdam, QFedAvg. Они решают проблемы гетерогенности данных и нестабильной сходимости.

Заключение и следующие шаги

Мы разобрали полный цикл федеративного обучения NLP-модели: от загрузки IMDB через datasets до запуска сервера Flower со стратегией FedAvg. Код клиента, сервера и подготовки данных воспроизводим на любом Linux-окружении с Python 3.10+.

Следующие шаги для углубления: добавьте дифференциальную приватность через flwr.client.DPClient, протестируйте FedProx на не-IID данных, масштабируйте на реальные устройства с помощью Flower Simulation Engine. Документация Flower и Hugging Face покрывает эти темы с примерами кода.

Если вы работаете с открытыми моделями и хотите лучше понимать инфраструктуру Hugging Face, посмотрите разбор huggingface_hub v1.0. Для тех, кто интересуется сжатием моделей перед развёртыванием на клиентах, полезно руководство по дистилляции LLM.

Подписаться на канал