Spec-Zone.ru › PyTorch 2.14

LinearCrossEntropyOptions

class torch.nn.LinearCrossEntropyOptions(allow_retain_graph=False, batch_chunk_size=None, chunking_method='auto', acc_policy='auto', acc_dtype=None) [исходный код]

Настройки для фрагментированной реализации linear_cross_entropy().

Фрагментированная реализация обрабатывает измерение пакета частями, поэтому полный тензор логитов (num_batches, num_classes) не материализуется — это полезно, когда num_classes намного больше, чем in_features (например, для голов выходного слоя со словарём LLM). Передайте options=None, чтобы использовать эталонный путь; передайте экземпляр этого класса, чтобы включить фрагментированную реализацию.

Конструктор LinearCrossEntropyOptions() без аргументов оставляет acc_policy и chunking_method равными "auto"; их значения определяются при вызове с учётом устройства и типа данных — см. описания полей ниже.

Поддерживается подмножество конфигураций linear_cross_entropy(); неподдерживаемые конфигурации переключаются на эталонный путь с предупреждением.

Фрагментирование эффективно, когда num_batches >= in_features и num_classes > in_features; при меньших значениях эталонный путь обходится дешевле.

acc_dtype: dtype | None

Тип данных для внутренних накоплений. None во время вызова принимает значение torch.float32 при acc_policy="auto" с входными данными fp16/bf16 на оборудовании с умножением матриц со смешанной точностью (CUDA SM 7.0+ для fp16, SM 8.0+ для bf16 и CPU); в остальных случаях используется тип данных входных данных. Для смешанной точности в настоящее время требуются входные данные fp16/bf16 с acc_dtype=torch.float32.

acc_policy: Literal['accurate', 'compact', 'auto']

Компромисс между точностью и использованием памяти для фрагментированного пути. Определяет, какие промежуточные результаты хранятся в acc_dtype, а какие — в типе данных входных данных, а также создаётся ли временный буфер градиента весов для каждого фрагмента.

  • "auto" (по умолчанию) — "accurate" для входных данных низкой точности на CPU (fp16/bf16), "compact" в остальных случаях (_auto_acc_policy()). Явно передайте "accurate" для любой другой платформы, эмулирующей GEMM fp16/bf16 с повышением точности до fp32.
  • "accurate" — наиболее широкое использование acc_dtype; заметно повышает точность градиента входных данных, когда размер фрагмента велик относительно num_classes. Наибольший пиковый расход памяти и самая низкая скорость среди фрагментированных политик на CUDA. Единственная фрагментированная политика, при которой умножение матриц для градиента весов выполняется в fp32 на CPU (при других политиках используется эмулируемый на CPU путь низкой точности, который примерно в 20–50 раз медленнее).
  • "compact" — использует acc_dtype только там, где это необходимо для корректности градиентов; накапливает градиент весов для каждого фрагмента напрямую с помощью addmm_ вместо временного буфера (num_classes, in_features) acc_dtype (на CUDA cuBLAS использует внутренний аккумулятор fp32, поэтому точность массовых вычислений не меняется). Экономит num_classes * in_features * sizeof(acc_dtype) — обычно несколько сотен МБ для голов выходного слоя LLM. При смешанной точности на платформах, отличных от CUDA, этот временный буфер сохраняется для межфрагментного накопления.

Разница в точности между "compact" и "accurate" заметна только тогда, когда acc_dtype отличается от типа данных входных данных; "compact" экономит память в обоих случаях.

allow_retain_graph: bool

Разрешает retain_graph=True при обратном проходе. Применяется только к скалярным операциям редукции ("mean" / "sum").

Если значение равно False (по умолчанию), обратный проход этих операций использует заранее вычисленные буферы градиентов на месте; повторный .backward() вызывает RuntimeError.

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

reduction="none" игнорирует это поле: при обратном проходе фрагментированные градиенты пересчитываются по сохранённым входным данным, поэтому retain_graph=True всегда работает без дополнительного выделения памяти.

Автоматическое дифференцирование высших порядков (gradgrad, AD в прямом режиме) не поддерживается.

При использовании torch.compile() для скалярных операций редукции это поле автоматически получает значение True, поскольку проверка повторного обратного прохода в режиме по умолчанию опирается на изменение ctx, которое Dynamo не сохраняет; оболочка выдаёт предупреждение об этом изменении.

batch_chunk_size: int | None

Количество строк пакета в одном фрагменте. Операция выполняет цикл по ceil(num_batches / batch_chunk_size) фрагментам; меньшие значения снижают пиковый расход памяти, но увеличивают число запусков ядер. Значение по умолчанию None означает один фрагмент. Нельзя использовать одновременно с chunking_method — если заданы оба параметра и их значения не совпадают, возникает ValueError.

chunking_method: str | None

Эвристика выбора значения batch_chunk_size.

  • "auto" (по умолчанию) — принимает значение "aspect_ratio" (коэффициент 1); для компактного пути дополнительно ограничивается значением B_ref на целевой элемент в _adjust(), чтобы пиковый расход памяти при фрагментировании не превышал расход без фрагментирования в бюджетном режиме (in_features >= num_classes).
  • "aspect_ratio" — выбирает размер каждого фрагмента так, чтобы его буфер логитов (batch_chunk_size, num_classes) занимал в памяти столько же, сколько входные данные (num_batches, in_features): next_pow2(ceil(num_batches / ceil(num_classes / in_features))). Оптимальный вариант, когда num_classes >> in_features (головы выходного слоя со словарём LLM).
  • "aspect_ratio:N" (N >= 1) — то же самое, но значение делится на N. Примерно в N раз меньше пиковый расход памяти ценой увеличения количества фрагментов в N раз.
  • None — отключает эвристику; используется значение batch_chunk_size.

© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/generated/torch.nn.LinearCrossEntropyOptions.html

Spec-Zone.ru

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