使用mint.index_select 在图模式下求梯度报错AssertionError
收藏回复举报
使用mint.index_select 在图模式下求梯度报错AssertionError
发表于2025-02-21 16:22:41
0 查看

1 系统环境

硬件环境(Ascend/GPU/CPU): Ascend/GPU/CPU

MindSpore版本: mindspore=2.4

执行模式(PyNative/ Graph): Graph

Python版本: Python=3.9

操作系统平台: linux

2 报错信息

2.1 问题描述

使用mint.index_select在图模式如果没有unset RANK_TABLE_FILE,就会报错RuntimeError

2.2 报错信息

不开启unet:

=============================================================== FAILURES ===============================================================
__________________________________________________ test_index_select_forward_back[0] ___________________________________________________
mode = 0

    @pytest.mark.parametrize('mode', [ms.GRAPH_MODE, ms.PYNATIVE_MODE])
    def test_index_select_forward_back(mode):
        """测试前向和反向传播,对比梯度"""
        ms.set_context(mode=mode)
   
        def forward_ms(x):
            return mint.index_select(x, 0, ms_index).sum()
   
        def forward_torch(x):
            return torch.index_select(x, 0, torch_index).sum()
   
        try:
            input_data = [[1, 6, 2, 4], [7, 3, 8, 2], [2, 9, 11, 5]]
            index = [0, 2]
            ms_tensor, torch_tensor, ms_index, torch_index = create_tensors(input_data, ms.float32, torch.float32, index=index, requires_grad=True)
            grad_fn_ms = value_and_grad(forward_ms)
            output_ms, gradient_ms = grad_fn_ms(ms_tensor)
   
            output_torch = forward_torch(torch_tensor)
            output_torch.backward()
            compare_results(output_ms, output_torch.detach())
            compare_results(gradient_ms, torch_tensor.grad)
        except Exception as e:
>           raise e

test_index_select.py:167:
_ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _
test_index_select.py:160: in test_index_select_forward_back
    output_ms, gradient_ms = grad_fn_ms(ms_tensor)
../anaconda3/envs/MindSpore/lib/python3.9/site-packages/mindspore/common/api.py:960: in staging_specialize
    out = _MindsporeFunctionExecutor(func, hash_obj, dyn_args, process_obj, jit_config)(*args, **kwargs)
../anaconda3/envs/MindSpore/lib/python3.9/site-packages/mindspore/common/api.py:188: in wrapper
    results = fn(*arg, **kwargs)
../anaconda3/envs/MindSpore/lib/python3.9/site-packages/mindspore/common/api.py:582: in __call__
    raise err
../anaconda3/envs/MindSpore/lib/python3.9/site-packages/mindspore/common/api.py:579: in __call__
    phase = self.compile(self.fn.__name__, *args_list, **kwargs)
_ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _
self = <mindspore.common.api._MindsporeFunctionExecutor object at 0xffff8a646b80>, method_name = 'after_grad'
args = (Tensor(shape=[3, 4], dtype=Float32, value=
[[ 1.00000000e+00,  6.00000000e+00,  2.00000000e+00,  4.00000000e+00],
 [ ...00000e+00,  8.00000000e+00,  2.00000000e+00],
 [ 2.00000000e+00,  9.00000000e+00,  1.10000000e+01,  5.00000000e+00]]),)
kwargs = {}
compile_args = (Tensor(shape=[3, 4], dtype=Float32, value=
[[ 1.00000000e+00,  6.00000000e+00,  2.00000000e+00,  4.00000000e+00],
 [ ...00000e+00,  8.00000000e+00,  2.00000000e+00],
 [ 2.00000000e+00,  9.00000000e+00,  1.10000000e+01,  5.00000000e+00]]),)
key_id = '1876501395697841734756182555320576'
generate_name = 'mindspore.ops.composite.base.after_grad./home/ma-user/anaconda3/envs/MindSpore/lib/python3.9/site-packages/mindspore/ops/composite/base.py.599.1734756182555320576'
echo_function_name = 'function "after_grad" at the file "/home/ma-user/anaconda3/envs/MindSpore/lib/python3.9/site-packages/mindspore/ops/composite/base.py", line 599'
full_function_name = 'mindspore.ops.composite.base.after_grad./home/ma-user/anaconda3/envs/MindSpore/lib/python3.9/site-packages/mindspore/ops/composite/base.py.599'
create_time = '1734756182555320576', key = 0, parameter_ids = ''
phase = 'mindspore.ops.composite.base.after_grad./home/ma-user/anaconda3/envs/MindSpore/lib/python3.9/site-packages/mindspore/ops/composite/base.py.599.1734756182555320576.0'
jit_config_dict = {'debug_level': 'RELEASE', 'exc_mode': 'auto', 'infer_boost': 'off', 'jit_level': '', ...}

    def compile(self, method_name, *args, **kwargs):
        """Returns pipeline for the given args."""
        # Check whether hook function registered on Cell object.
        if self.obj and hasattr(self.obj, "_hook_fn_registered"):
            if self.obj._hook_fn_registered():
                logger.warning(f"For 'Cell', it's not support hook function when using 'jit' decorator. "
                               f"If you want to use hook function, please use context.set_context to set "
                               f"pynative mode and remove 'jit' decorator.")
        # Chose dynamic shape tensors or actual input tensors as compile args.
        compile_args = self._generate_compile_args(args)
        key_id = self._get_key_id()
        compile_args = get_auto_dynamic_shape_args_with_check_input_signature(compile_args, key_id,
                                                                              self.input_signature)
   
        # Restore the mutable attr for every arg.
        compile_args = _restore_mutable_attr(args, compile_args)
        self._compile_args = compile_args
        generate_name, echo_function_name = self._get_generate_name()
        # The full Function name
        full_function_name = generate_name
        create_time = ''
   
        # Add key with obj
        if self.obj is not None:
            if self.obj.__module__ != self.fn.__module__:
                logger.info(
                    f'The module of `self.obj`: `{self.obj.__module__}` is not same with the module of `self.fn`: '
                    f'`{self.fn.__module__}`')
            self.obj.__parse_method__ = method_name
            if isinstance(self.obj, ms.nn.Cell):
                generate_name = generate_name + '.' + str(self.obj.create_time)
                create_time = str(self.obj.create_time)
            else:
                generate_name = generate_name + '.' + str(self._create_time)
                create_time = str(self._create_time)
   
            generate_name = generate_name + '.' + str(id(self.obj))
            full_function_name = generate_name
        else:
            # Different instance of same class may use same memory(means same obj_id) at diff times.
            # To avoid unexpected phase matched, add create_time to generate_name.
            generate_name = generate_name + '.' + str(self._create_time)
            create_time = str(self._create_time)
   
        self.enable_tuple_broaden = False
        if hasattr(self.obj, "enable_tuple_broaden"):
            self.enable_tuple_broaden = self.obj.enable_tuple_broaden
   
        self._graph_executor.set_enable_tuple_broaden(self.enable_tuple_broaden)
        key = self._graph_executor.generate_arguments_key(self.fn, compile_args, kwargs, self.enable_tuple_broaden)
    
        parameter_ids = _get_parameter_ids(args, kwargs)
        if parameter_ids != "":
            key = str(key) + '.' + parameter_ids
        phase = generate_name + '.' + str(key)
   
        update_auto_dynamic_shape_phase_with_check_input_signature(compile_args, key_id, phase, self.input_signature)
   
        if phase in ms_compile_cache:
            # Release resource should be released when CompileInner won't be executed, such as cur_convert_input_
            # generated in generate_arguments_key.
            self._graph_executor.clear_compile_arguments_resource()
            return phase
   
        _check_recompile(self.obj, compile_args, kwargs, full_function_name, create_time, echo_function_name)
   
        # If enable compile cache, get the dependency files list and set to graph executor.
        self._set_compile_cache_dep_files()
        if self.jit_config_dict:
            self._graph_executor.set_jit_config(self.jit_config_dict)
        else:
            jit_config_dict = JitConfig().jit_config_dict
            self._graph_executor.set_jit_config(jit_config_dict)

        if self.obj is None:
            # Set an attribute to fn as an identifier.
            if isinstance(self.fn, types.MethodType):
                setattr(self.fn.__func__, "__jit_function__", True)
            else:
                setattr(self.fn, "__jit_function__", True)
>           is_compile = self._graph_executor.compile(self.fn, compile_args, kwargs, phase, True)
E           RuntimeError: Compile graph kernel_graph0 failed.
E          
E           ----------------------------------------------------
E           - C++ Call Stack: (For framework developers)
E           ----------------------------------------------------
E           mindspore/ccsrc/plugin/device/ascend/hal/hardware/ge_graph_executor.cc:636 CompileGraph
../anaconda3/envs/MindSpore/lib/python3.9/site-packages/mindspore/common/api.py:674: RuntimeError

开启:

test_index_select.py:165: in test_index_select_forward_back
    compare_results(gradient_ms, torch_tensor.grad)_ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _
ms_result = Tensor(shape=[3, 4], dtype=Float32, value=
[[-9.15310234e-02,  1.28829825e+00, -1.13139284e+00, -2.20877814e+00],
 [-7...1860191e-01,  6.35216951e-01,  2.55284572e+00],
 [-9.15310234e-02,  1.28829825e+00, -1.13139284e+00, -2.20877814e+00]])
torch_result = tensor([[1., 1., 1., 1.],
        [0., 0., 0., 0.],
        [1., 1., 1., 1.]]), atol = 0.001
    def compare_results(ms_result, torch_result, atol=1e-3):
>       assert np.allclose(ms_result.asnumpy(), torch_result.numpy(), atol=atol)
E       assert False
E        +  where False = <function allclose at 0xffff078191f0>(array([[-0.09153102,  1.2882982 , -1.1313928 , -2.2087781 ],\n       [-0.75677174, -0.4818602 ,  0.63521695,  2.5528457 ],\n       [-0.09153102,  1.2882982 , -1.1313928 , -2.2087781 ]],\n      dtype=float32), array([[1., 1., 1., 1.],\n       [0., 0., 0., 0.],\n       [1., 1., 1., 1.]], dtype=float32), atol=0.001)
E        +    where <function allclose at 0xffff078191f0> = np.allclose
E        +    and   array([[-0.09153102,  1.2882982 , -1.1313928 , -2.2087781 ],\n       [-0.75677174, -0.4818602 ,  0.63521695,  2.5528457 ],\n       [-0.09153102,  1.2882982 , -1.1313928 , -2.2087781 ]],\n      dtype=float32) = <bound method Tensor.asnumpy of Tensor(shape=[3, 4], dtype=Float32, value=\n[[-9.15310234e-02,  1.28829825e+00, -1.1313...860191e-01,  6.35216951e-01,  2.55284572e+00],\n [-9.15310234e-02,  1.28829825e+00, -1.13139284e+00, -2.20877814e+00]])>()
E        +      where <bound method Tensor.asnumpy of Tensor(shape=[3, 4], dtype=Float32, value=\n[[-9.15310234e-02,  1.28829825e+00, -1.1313...860191e-01,  6.35216951e-01,  2.55284572e+00],\n [-9.15310234e-02,  1.28829825e+00, -1.13139284e+00, -2.20877814e+00]])> = Tensor(shape=[3, 4], dtype=Float32, value=\n[[-9.15310234e-02,  1.28829825e+00, -1.13139284e+00, -2.20877814e+00],\n [-7...1860191e-01,  6.35216951e-01,  2.55284572e+00],\n [-9.15310234e-02,  1.28829825e+00, -1.13139284e+00, -2.20877814e+00]]).asnumpy
E        +    and   array([[1., 1., 1., 1.],\n       [0., 0., 0., 0.],\n       [1., 1., 1., 1.]], dtype=float32) = <built-in method numpy of Tensor object at 0xffffa45d3040>()
E        +      where <built-in method numpy of Tensor object at 0xffffa45d3040> = tensor([[1., 1., 1., 1.],\n        [0., 0., 0., 0.],\n        [1., 1., 1., 1.]]).numpy
test_index_select.py:28: AssertionError

2.3 脚本信息

在终端分别使用或者不使用unset RANK_TABLE_FILE,再使用下面的脚本。

import pytest
import numpy as np
import mindspore as ms
from mindspore import Tensor, value_and_grad, mint
import torch

def create_tensors(input_data, ms_dtype, torch_dtype, index=None, requires_grad=False):
    ms_tensor = Tensor(input_data, ms_dtype)
    torch_tensor = torch.tensor(input_data, dtype=torch_dtype, requires_grad=requires_grad)
    if index is not None:
        ms_index = Tensor(index, ms.int32)
        torch_index = torch.tensor(index, dtype=torch.long)
    else:
        ms_index = None
        torch_index = None
    return ms_tensor, torch_tensor, ms_index, torch_index

def perform_index_select(ms_tensor, torch_tensor, dim, ms_index, torch_index):
    ms_result = mint.index_select(ms_tensor, dim, ms_index)
    torch_result = torch.index_select(torch_tensor, dim, torch_index)
    return ms_result, torch_result

def compare_results(ms_result, torch_result, atol=1e-3):
    assert np.allclose(ms_result.asnumpy(), torch_result.numpy(), atol=atol)

@pytest.mark.parametrize('mode', [ms.GRAPH_MODE, ms.PYNATIVE_MODE])
def test_index_select_forward_back(mode):
    """测试前向和反向传播,对比梯度"""
    ms.set_context(mode=mode)
    def forward_ms(x):
        return mint.index_select(x, 0, ms_index).sum()
    def forward_torch(x):
        return torch.index_select(x, 0, torch_index).sum()

    try:
        input_data = [[1, 6, 2, 4], [7, 3, 8, 2], [2, 9, 11, 5]]
        index = [0, 2]
        ms_tensor, torch_tensor, ms_index, torch_index = create_tensors(input_data, ms.float32, torch.float32, index=index, requires_grad=True)
        grad_fn_ms = value_and_grad(forward_ms)
        output_ms, gradient_ms = grad_fn_ms(ms_tensor)
        output_torch = forward_torch(torch_tensor)
        output_torch.backward()
        compare_results(output_ms, output_torch.detach())
        compare_results(gradient_ms, torch_tensor.grad)
    except Exception as e:
        raise e

3 根因分析

******此处由用户填写******

4 解决方案

******此处由用户填写******

包含文字方案和最终脚本代码

请将正确的脚本打包并上传附件

我要发帖子