Spec-Zone.ru › PyTorch 2.14

torch.fx.passes.reinplace.reinplace

torch.fx.passes.reinplace.reinplace(gm, *sample_args) [исходный код]

Принимает fx.GraphModule и модифицирует его, выполняя «преобразование в операции на месте» и изменяя узлы графа. Мы ищем места вызова операций не на месте, например b = a.add(...), и преобразуем их в операции на месте (b = a.add_(...)), если вход текущего оператора («a») нигде далее в графе не используется повторно.

В настоящее время этот проход рассчитан на работу с функциональным графом ATen. Такой граф можно получить, запустив make_fx(functionalize(f)).

Для определения отношений алиасинга входных данных нужны примеры входных данных. В общем случае нельзя преобразовать узел b = a.add(...) в операцию на месте, если «a» имеет общий алиас с любым из входов программы.

Для заданного узла b = foo(a, args...) алгоритм преобразования в операцию на месте выглядит следующим образом:

(1) Выполнить начальные проверки метаданных «a» и «args…», которые могут исключить возможность преобразования в операцию на месте.

  • (1a) Проверить, что аргумент self, который мы пытаемся преобразовать в операцию на месте, имеет допустимые метаданные типа данных и размера.

    Например, если у нас есть:

    a = torch.ones(1)
    b = torch.ones(10)
    out = torch.add(a, b)
    

    Мы не можем преобразовать это в a.add_(b), поскольку для этого пришлось бы изменить размер «a».

    Аналогично, мы не можем преобразовать torch.ge(a, b) в a.ge_(b), поскольку для этого пришлось бы изменить тип данных «a» (например, с float32 на bool). Обратите внимание, что в этом конкретном примере технически можно было бы сделать лучше...

    Если мы видим такой шаблон:

    a_1 = a.ge(b)
    a_2 = aten._to_copy(a_1, a.dtype)
    

    то его можно полностью преобразовать в операции на месте (именно такой код выдаёт functionalization, когда обнаруживает a.ge_(b)).

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

    Также нельзя преобразовывать в операции на месте тензоры с перекрывающейся памятью, например torch.ones(1).expand(4, 4).add_(1).

  • (1b) Проверить, является ли «a» алиасом какого-либо входа программы.

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

    ПРИМЕЧАНИЕ: в будущем следует добавить такую оптимизацию: если «a» является входом программы (или его алиасом), но далее в программе есть узел вида a.copy_(...), то преобразование в операцию на месте допустимо — мы временно повторно используем буфер «a», который впоследствии будет перезаписан вызовом copy_().

    Эта оптимизация будет важна для программ, изменяющих свои входные данные. Пока она не реализована.

  • (1c) Проверить, имеют ли «a» и «args…» общий алиас.

    Например, преобразование в операцию на месте с созданием кода, подобного приведённому ниже, не гарантированно будет корректным:

    aten.mul_(a, a)
    

(2) Проверить, что «a» и все его существующие алиасы нигде далее в графе не используются. Если это так, безопасно преобразовать b = foo_(a) в операцию на месте.

Здесь есть несколько оговорок, подробнее описанных ниже:

  • (a) Если далее «a» используется как аргумент операции-представления, это допустимо. Проблема возникает только в том случае, если «a» (или это представление) далее передаётся обычному оператору либо возвращается в качестве результата программы.
  • (b) Если «a» является повторяющимся аргументом в foo(), не выполняйте преобразование в операцию на месте. Большинство ядер ATen не гарантируют корректность такого случая, например, если выполнить aten.mul_(a, a). Поэтому в этом случае мы просто запрещаем преобразование в операцию на месте.
  • (c) Если «a» используется как вход операции «обратного преобразования» представления / scatter, преобразование в операцию на месте потенциально допустимо (и эту операцию scatter можно удалить из графа). Более подробный пример приведён ниже.

ПРИМЕЧАНИЕ: на этом этапе выполняется оптимизация, необходимая для полного восстановления производительности после functionalization.

Рассмотрим такую программу:

def f(x):
    a = torch.ops.aten.add(x, x)
    b = torch.ops.aten.diagonal(a)
    torch.ops.aten.fill_(b, 0)
    return d

Functionalization выдаст следующий код:

def f(x):
    a = torch.ops.aten.add(x, x)
    b = torch.ops.aten.diagonal(a, 0, 1)
    b_updated = torch.ops.aten.fill(b, 0)
    a_updated = torch.ops.aten.diagonal_scatter(a, b_updated, 0, 1)
    return a_updated

Обычно мы не смогли бы преобразовать fill в операцию на месте, поскольку «b» имеет общий алиас с «a», который используется вызовом diagonal_scatter.

При «преобразовании в операции на месте» нужно определить, что вызов diagonal_scatter можно полностью удалить, если преобразовать add() в операцию на месте.

Поэтому для каждого alias in alias_set(a) вместо проверки того, что «alias» нигде далее в графе не используется, мы проверяем, что выполняется ОДНО из условий:

    1. alias нигде далее в графе не используется; ИЛИ
  • (b) alias используется далее в графе ровно один раз — в следующей операции:

    out = foo_scatter(alias, x, args...)
    

    при этом должны выполняться следующие условия:

    • (i) foo_scatter — оператор «обратного преобразования» для foo. Это относится только к операциям foo, представляющим собой операции-представления, которые создают представление подмножества памяти исходного тензора. На практике таких операций примерно четыре:

      diagonal -> diagonal_scatter
      slice -> slice_scatter
      select -> select_scatter
      as_strided -> as_strided_scatter
      
    • (ii) «args…» должны совпадать в вызовах foo() и foo_scatter().

(3) Выполнить фактическое преобразование foo в операцию на месте!

Случай (3b) является типичным, но для {view}_scatter (3a) требуется особая обработка.

  • (3a) Операции {view}_scatter.

    Рассмотрим такую программу:

    a = torch.zeros(2, 2)
    b = torch.ones(2)
    a[0] = b
    

    После functionalization она будет выглядеть так:

    a = torch.zeros(2)
    b = torch.ones(1)
    a_updated = torch.select_scatter(a, b, 0, 0)
    

    Однако в этом случае нет «функциональной» операции, которую можно преобразовать в операцию на месте! Вместо этого мы хотим напрямую удалить вызов select_scatter. Из шага (3) мы уже знаем, что это допустимо, поскольку у «a» нет дальнейших использований в графе.

    Мы преобразуем операцию {view}_scatter в операцию на месте следующим образом.

    До:

    a_updated = torch.select_scatter(a, b, args...)
    

    После:

    a_slice = a.select(a, args...)
    a_slice.copy_(b)
    
  • (3b) В противном случае заменить функциональную операцию её вариантом на месте.

    До:

    b = foo(a, args...)
    

    После:

    a.foo_(args...)
    

(4) Наконец, после преобразования одного из вариантов:

# Before:                              # After:
b = foo(a)                             foo_(a)

или:

# Before:
b = {slice}_scatter(a, mutated_slice, args...)
# After:
slice = {slice}(a, args...)
slice.copy_(mutated_slice)

Нужно найти все последующие узлы, использующие «b» в качестве аргумента, и заменить его на «a».

Обратите внимание, что для большинства операций на месте это не требуется (поскольку большинство таких операций возвращают «self» в качестве результата). Однако это верно не для всех изменяющих состояние операций, поэтому нам действительно нужно заменить все аргументы.

Нужно также обновить метаданные Dict[StorageWeakRef, Set[Node]], которые сопоставляют хранилищу тензора множество всех узлов, использующих это хранилище в качестве входных данных. В частности, преобразование b = foo(a) в операцию на месте приводит к объединению множеств «a» и «b».

(5) Все узлы view_inverse/scatter, которые на шаге (3) были отмечены как «их можно игнорировать», удаляются из графа вручную. Их результаты больше нигде не используются, поэтому стандартный DCE мог бы удалить их автоматически, однако теперь мы не можем запускать проход DCE из FX, поскольку в графе появились изменяющие состояние операции.

Примечание

Обратная совместимость этого API гарантируется.

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

GraphModule

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

Spec-Zone.ru

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