已合并
[需求] ACLGraph权重加载:linear+loader #6
huanglei创建于 2025年12月30日
[需求] ACLGraph权重加载:linear+loader #6
已合并
从已删除 :br_linear合入到Katrina-CXY/MindIE-LLM_opensourceaclgraph_final
共 15 个文件变更+4312-0
The file is empty
| @@ -0,0 +1,424 @@ | |||
| 1 | +# Copyright (c) Huawei Technologies Co., Ltd. 2025-2026. All rights reserved. | ||
| 2 | +# MindIE is licensed under Mulan PSL v2. | ||
| 3 | +# You can use this software according to the terms and conditions of the Mulan PSL v2. | ||
| 4 | +# You may obtain a copy of Mulan PSL v2 at: | ||
| 5 | +# http://license.coscl.org.cn/MulanPSL2 | ||
| 6 | +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, | ||
| 7 | +# EITHER EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, | ||
| 8 | +# MERCHANTABILITY OR FIT FOR A PARTICULAR PURPOSE. | ||
| 9 | +# See the Mulan PSL v2 for more details. | ||
| 10 | + | ||
| 11 | +import unittest | ||
| 12 | +from unittest.mock import MagicMock, patch | ||
| 13 | +import torch | ||
| 14 | + | ||
| 15 | +from mindie_llm.runtime.layers.quantization.ms_model_slim.anti_outlier import AntiOutlierNormMethod | ||
| 16 | +from mindie_llm.runtime.layers.parameter import BaseParameter | ||
| 17 | + | ||
| 18 | + | ||
| 19 | +class TestAntiOutlierNormMethod(unittest.TestCase): | ||
| 20 | + """Test cases for AntiOutlierNormMethod.""" | ||
| 21 | + | ||
| 22 | + def setUp(self): | ||
| 23 | + """Set up test fixtures.""" | ||
| 24 | + self.quant_method = AntiOutlierNormMethod() | ||
| 25 | + self.hidden_size = 512 | ||
| 26 | + | ||
| 27 | + def test_create_weights(self): | ||
| 28 | + """Test create_weights method.""" | ||
| 29 | + # Create a mock layer | ||
| 30 | + mock_layer = MagicMock(spec=torch.nn.Module) | ||
| 31 | + mock_layer.register_parameter = MagicMock() | ||
| 32 | + | ||
| 33 | + params_dtype = torch.float32 | ||
| 34 | + extra_attrs = {"output_dim": 0} | ||
| 35 | + | ||
| 36 | + self.quant_method.create_weights( | ||
| 37 | + layer=mock_layer, | ||
| 38 | + hidden_size=self.hidden_size, | ||
| 39 | + params_dtype=params_dtype, | ||
| 40 | + **extra_attrs, | ||
| 41 | + ) | ||
| 42 | + | ||
| 43 | + # Verify register_parameter was called twice (weight and bias) | ||
| 44 | + self.assertEqual(mock_layer.register_parameter.call_count, 2) | ||
| 45 | + | ||
| 46 | + # Get the registered parameters | ||
| 47 | + weight_call = mock_layer.register_parameter.call_args_list[0] | ||
| 48 | + bias_call = mock_layer.register_parameter.call_args_list[1] | ||
| 49 | + | ||
| 50 | + # Verify weight parameter | ||
| 51 | + self.assertEqual(weight_call[0][0], "weight") | ||
| 52 | + weight_param = weight_call[0][1] | ||
| 53 | + self.assertIsInstance(weight_param, BaseParameter) | ||
| 54 | + self.assertEqual(weight_param.data.shape, (self.hidden_size,)) | ||
| 55 | + self.assertEqual(weight_param.data.dtype, params_dtype) | ||
| 56 | + # Verify weight is initialized to ones | ||
| 57 | + self.assertTrue(torch.allclose(weight_param.data, torch.ones(self.hidden_size, dtype=params_dtype))) | ||
| 58 | + | ||
| 59 | + # Verify bias parameter | ||
| 60 | + self.assertEqual(bias_call[0][0], "bias") | ||
| 61 | + bias_param = bias_call[0][1] | ||
| 62 | + self.assertIsInstance(bias_param, BaseParameter) | ||
| 63 | + self.assertEqual(bias_param.data.shape, (self.hidden_size,)) | ||
| 64 | + self.assertEqual(bias_param.data.dtype, params_dtype) | ||
| 65 | + # Verify bias is initialized to zeros | ||
| 66 | + self.assertTrue(torch.allclose(bias_param.data, torch.zeros(self.hidden_size, dtype=params_dtype))) | ||
| 67 | + | ||
| 68 | + def test_create_weights_with_custom_dtype(self): | ||
| 69 | + """Test create_weights with custom dtype.""" | ||
| 70 | + mock_layer = MagicMock(spec=torch.nn.Module) | ||
| 71 | + mock_layer.register_parameter = MagicMock() | ||
| 72 | + | ||
| 73 | + params_dtype = torch.float16 | ||
| 74 | + | ||
| 75 | + self.quant_method.create_weights( | ||
| 76 | + layer=mock_layer, | ||
| 77 | + hidden_size=self.hidden_size, | ||
| 78 | + params_dtype=params_dtype, | ||
| 79 | + ) | ||
| 80 | + | ||
| 81 | + # Verify weight dtype | ||
| 82 | + weight_call = mock_layer.register_parameter.call_args_list[0] | ||
| 83 | + weight_param = weight_call[0][1] | ||
| 84 | + self.assertEqual(weight_param.data.dtype, params_dtype) | ||
| 85 | + | ||
| 86 | + # Verify bias dtype | ||
| 87 | + bias_call = mock_layer.register_parameter.call_args_list[1] | ||
| 88 | + bias_param = bias_call[0][1] | ||
| 89 | + self.assertEqual(bias_param.data.dtype, params_dtype) | ||
| 90 | + | ||
| 91 | + def test_create_weights_with_extra_attrs(self): | ||
| 92 | + """Test create_weights with extra attributes.""" | ||
| 93 | + mock_layer = MagicMock(spec=torch.nn.Module) | ||
| 94 | + mock_layer.register_parameter = MagicMock() | ||
| 95 | + | ||
| 96 | + extra_attrs = {"output_dim": 0, "input_dim": 1, "custom_attr": "test"} | ||
| 97 | + | ||
| 98 | + self.quant_method.create_weights( | ||
| 99 | + layer=mock_layer, | ||
| 100 | + hidden_size=self.hidden_size, | ||
| 101 | + params_dtype=torch.float32, | ||
| 102 | + **extra_attrs, | ||
| 103 | + ) | ||
| 104 | + | ||
| 105 | + # Verify extra attributes were added to weight | ||
| 106 | + weight_call = mock_layer.register_parameter.call_args_list[0] | ||
| 107 | + weight_param = weight_call[0][1] | ||
| 108 | + self.assertEqual(getattr(weight_param,"output_dim"), 0) | ||
| 109 | + self.assertEqual(getattr(weight_param,"input_dim"), 1) | ||
| 110 | + self.assertEqual(getattr(weight_param,"custom_attr"), "test") | ||
| 111 | + | ||
| 112 | + # Verify extra attributes were added to bias | ||
| 113 | + bias_call = mock_layer.register_parameter.call_args_list[1] | ||
| 114 | + bias_param = bias_call[0][1] | ||
| 115 | + self.assertEqual(getattr(bias_param,"output_dim"), 0) | ||
| 116 | + self.assertEqual(getattr(bias_param,"input_dim"), 1) | ||
| 117 | + self.assertEqual(getattr(bias_param,"custom_attr"), "test") | ||
| 118 | + | ||
| 119 | + def test_create_weights_different_hidden_sizes(self): | ||
| 120 | + """Test create_weights with different hidden sizes.""" | ||
| 121 | + mock_layer = MagicMock(spec=torch.nn.Module) | ||
| 122 | + mock_layer.register_parameter = MagicMock() | ||
| 123 | + | ||
| 124 | + for hidden_size in [128, 256, 512, 1024, 2048]: | ||
| 125 | + mock_layer.register_parameter.reset_mock() | ||
| 126 | + | ||
| 127 | + self.quant_method.create_weights( | ||
| 128 | + layer=mock_layer, | ||
| 129 | + hidden_size=hidden_size, | ||
| 130 | + params_dtype=torch.float32, | ||
| 131 | + ) | ||
| 132 | + | ||
| 133 | + # Verify weight shape | ||
| 134 | + weight_call = mock_layer.register_parameter.call_args_list[0] | ||
| 135 | + weight_param = weight_call[0][1] | ||
| 136 | + self.assertEqual(weight_param.data.shape, (hidden_size,)) | ||
| 137 | + | ||
| 138 | + # Verify bias shape | ||
| 139 | + bias_call = mock_layer.register_parameter.call_args_list[1] | ||
| 140 | + bias_param = bias_call[0][1] | ||
| 141 | + self.assertEqual(bias_param.data.shape, (hidden_size,)) | ||
| 142 | + | ||
| 143 | + | ||
| 144 | + def test_apply_without_residual(self, mock_npu_rms_norm): | ||
| 145 | + """Test apply method without residual.""" | ||
| 146 | + # Create a mock layer | ||
| 147 | + mock_layer = MagicMock(spec=torch.nn.Module) | ||
| 148 | + mock_layer.weight = MagicMock() | ||
| 149 | + mock_layer.weight.data = torch.ones(self.hidden_size, dtype=torch.float32) | ||
| 150 | + mock_layer.bias = MagicMock() | ||
| 151 | + mock_layer.bias.data = torch.zeros(self.hidden_size, dtype=torch.float32) | ||
| 152 | + mock_layer.variance_epsilon = 1e-6 | ||
| 153 | + | ||
| 154 | + # Create input tensor | ||
| 155 | + x = torch.randn(2, 3, self.hidden_size, dtype=torch.float32) | ||
| 156 | + | ||
| 157 | + # Mock npu_rms_norm to return normalized tensor and variance | ||
| 158 | + normalized_tensor = torch.randn(2, 3, self.hidden_size, dtype=torch.float32) | ||
| 159 | + variance = torch.randn(2, 3, dtype=torch.float32) | ||
| 160 | + mock_npu_rms_norm.return_value = (normalized_tensor, variance) | ||
| 161 | + | ||
| 162 | + # Call apply | ||
| 163 | + output = self.quant_method.apply(layer=mock_layer, x=x) | ||
| 164 | + | ||
| 165 | + # Verify npu_rms_norm was called with correct arguments | ||
| 166 | + mock_npu_rms_norm.assert_called_once() | ||
| 167 | + call_args = mock_npu_rms_norm.call_args | ||
| 168 | + self.assertTrue(torch.equal(call_args[0][0], x)) | ||
| 169 | + self.assertTrue(torch.equal(call_args[0][1], mock_layer.weight.data)) | ||
| 170 | + self.assertEqual(call_args[0][2], mock_layer.variance_epsilon) | ||
| 171 | + | ||
| 172 | + # Verify output shape (should be normalized + bias) | ||
| 173 | + self.assertEqual(output.shape, (2, 3, self.hidden_size)) | ||
| 174 | + # Output should be normalized_tensor + bias | ||
| 175 | + expected_output = normalized_tensor + mock_layer.bias.data | ||
| 176 | + self.assertTrue(torch.equal(output, expected_output)) | ||
| 177 | + | ||
| 178 | + | ||
| 179 | + def test_apply_without_residual_different_shapes(self, mock_npu_rms_norm): | ||
| 180 | + """Test apply method without residual with different input shapes.""" | ||
| 181 | + mock_layer = MagicMock(spec=torch.nn.Module) | ||
| 182 | + mock_layer.weight = MagicMock() | ||
| 183 | + mock_layer.weight.data = torch.ones(self.hidden_size, dtype=torch.float32) | ||
| 184 | + mock_layer.bias = MagicMock() | ||
| 185 | + mock_layer.bias.data = torch.zeros(self.hidden_size, dtype=torch.float32) | ||
| 186 | + mock_layer.variance_epsilon = 1e-6 | ||
| 187 | + | ||
| 188 | + test_shapes = [ | ||
| 189 | + (self.hidden_size,), # 1D | ||
| 190 | + (10, self.hidden_size), # 2D | ||
| 191 | + (2, 3, self.hidden_size), # 3D | ||
| 192 | + (1, 2, 3, self.hidden_size), # 4D | ||
| 193 | + ] | ||
| 194 | + | ||
| 195 | + for shape in test_shapes: | ||
| 196 | + mock_npu_rms_norm.reset_mock() | ||
| 197 | + | ||
| 198 | + x = torch.randn(*shape, dtype=torch.float32) | ||
| 199 | + normalized_tensor = torch.randn(*shape, dtype=torch.float32) | ||
| 200 | + variance = torch.randn(shape[:-1], dtype=torch.float32) | ||
| 201 | + mock_npu_rms_norm.return_value = (normalized_tensor, variance) | ||
| 202 | + | ||
| 203 | + output = self.quant_method.apply(layer=mock_layer, x=x) | ||
| 204 | + | ||
| 205 | + # Verify output shape matches input shape | ||
| 206 | + self.assertEqual(output.shape, shape) | ||
| 207 | + mock_npu_rms_norm.assert_called_once() | ||
| 208 | + | ||
| 209 | + | ||
| 210 | + def test_apply_with_residual(self, mock_npu_add_rms_norm): | ||
| 211 | + """Test apply method with residual.""" | ||
| 212 | + # Create a mock layer | ||
| 213 | + mock_layer = MagicMock(spec=torch.nn.Module) | ||
| 214 | + mock_layer.weight = MagicMock() | ||
| 215 | + mock_layer.weight.data = torch.ones(self.hidden_size, dtype=torch.float32) | ||
| 216 | + mock_layer.bias = MagicMock() | ||
| 217 | + mock_layer.bias.data = torch.zeros(self.hidden_size, dtype=torch.float32) | ||
| 218 | + mock_layer.variance_epsilon = 1e-6 | ||
| 219 | + | ||
| 220 | + # Create input and residual tensors | ||
| 221 | + x = torch.randn(2, 3, self.hidden_size, dtype=torch.float32) | ||
| 222 | + residual = torch.randn(2, 3, self.hidden_size, dtype=torch.float32) | ||
| 223 | + | ||
| 224 | + # Mock npu_add_rms_norm to return normalized tensor, variance, and updated residual | ||
| 225 | + normalized_tensor = torch.randn(2, 3, self.hidden_size, dtype=torch.float32) | ||
| 226 | + variance = torch.randn(2, 3, dtype=torch.float32) | ||
| 227 | + updated_residual = torch.randn(2, 3, self.hidden_size, dtype=torch.float32) | ||
| 228 | + mock_npu_add_rms_norm.return_value = (normalized_tensor, variance, updated_residual) | ||
| 229 | + | ||
| 230 | + # Call apply | ||
| 231 | + output, output_residual = self.quant_method.apply(layer=mock_layer, x=x, residual=residual) | ||
| 232 | + | ||
| 233 | + # Verify npu_add_rms_norm was called with correct arguments | ||
| 234 | + mock_npu_add_rms_norm.assert_called_once() | ||
| 235 | + call_args = mock_npu_add_rms_norm.call_args | ||
| 236 | + self.assertTrue(torch.equal(call_args[0][0], x)) | ||
| 237 | + self.assertTrue(torch.equal(call_args[0][1], residual)) | ||
| 238 | + self.assertTrue(torch.equal(call_args[0][2], mock_layer.weight.data)) | ||
| 239 | + self.assertEqual(call_args[0][3], mock_layer.variance_epsilon) | ||
| 240 | + | ||
| 241 | + # Verify output shape (should be normalized + bias) | ||
| 242 | + self.assertEqual(output.shape, (2, 3, self.hidden_size)) | ||
| 243 | + # Output should be normalized_tensor + bias | ||
| 244 | + expected_output = normalized_tensor + mock_layer.bias.data | ||
| 245 | + self.assertTrue(torch.equal(output, expected_output)) | ||
| 246 | + | ||
| 247 | + # Verify residual is returned | ||
| 248 | + self.assertEqual(output_residual.shape, (2, 3, self.hidden_size)) | ||
| 249 | + self.assertTrue(torch.equal(output_residual, updated_residual)) | ||
| 250 | + | ||
| 251 | + | ||
| 252 | + def test_apply_with_residual_different_shapes(self, mock_npu_add_rms_norm): | ||
| 253 | + """Test apply method with residual with different input shapes.""" | ||
| 254 | + mock_layer = MagicMock(spec=torch.nn.Module) | ||
| 255 | + mock_layer.weight = MagicMock() | ||
| 256 | + mock_layer.weight.data = torch.ones(self.hidden_size, dtype=torch.float32) | ||
| 257 | + mock_layer.bias = MagicMock() | ||
| 258 | + mock_layer.bias.data = torch.zeros(self.hidden_size, dtype=torch.float32) | ||
| 259 | + mock_layer.variance_epsilon = 1e-6 | ||
| 260 | + | ||
| 261 | + test_shapes = [ | ||
| 262 | + (self.hidden_size,), # 1D | ||
| 263 | + (10, self.hidden_size), # 2D | ||
| 264 | + (2, 3, self.hidden_size), # 3D | ||
| 265 | + ] | ||
| 266 | + | ||
| 267 | + for shape in test_shapes: | ||
| 268 | + mock_npu_add_rms_norm.reset_mock() | ||
| 269 | + | ||
| 270 | + x = torch.randn(*shape, dtype=torch.float32) | ||
| 271 | + residual = torch.randn(*shape, dtype=torch.float32) | ||
| 272 | + normalized_tensor = torch.randn(*shape, dtype=torch.float32) | ||
| 273 | + variance = torch.randn(shape[:-1], dtype=torch.float32) | ||
| 274 | + updated_residual = torch.randn(*shape, dtype=torch.float32) | ||
| 275 | + mock_npu_add_rms_norm.return_value = (normalized_tensor, variance, updated_residual) | ||
| 276 | + | ||
| 277 | + output, output_residual = self.quant_method.apply(layer=mock_layer, x=x, residual=residual) | ||
| 278 | + | ||
| 279 | + # Verify output shape matches input shape | ||
| 280 | + self.assertEqual(output.shape, shape) | ||
| 281 | + self.assertEqual(output_residual.shape, shape) | ||
| 282 | + mock_npu_add_rms_norm.assert_called_once() | ||
| 283 | + | ||
| 284 | + | ||
| 285 | + def test_apply_without_residual_bias_addition(self, mock_npu_rms_norm): | ||
| 286 | + """Test that bias is added correctly when no residual.""" | ||
| 287 | + mock_layer = MagicMock(spec=torch.nn.Module) | ||
| 288 | + mock_layer.weight = MagicMock() | ||
| 289 | + mock_layer.weight.data = torch.ones(self.hidden_size, dtype=torch.float32) | ||
| 290 | + mock_layer.bias = MagicMock() | ||
| 291 | + # Set bias to non-zero values | ||
| 292 | + mock_layer.bias.data = torch.ones(self.hidden_size, dtype=torch.float32) * 0.5 | ||
| 293 | + mock_layer.variance_epsilon = 1e-6 | ||
| 294 | + | ||
| 295 | + x = torch.randn(2, 3, self.hidden_size, dtype=torch.float32) | ||
| 296 | + normalized_tensor = torch.randn(2, 3, self.hidden_size, dtype=torch.float32) | ||
| 297 | + variance = torch.randn(2, 3, dtype=torch.float32) | ||
| 298 | + mock_npu_rms_norm.return_value = (normalized_tensor, variance) | ||
| 299 | + | ||
| 300 | + output = self.quant_method.apply(layer=mock_layer, x=x) | ||
| 301 | + | ||
| 302 | + # Verify bias was added | ||
| 303 | + expected_output = normalized_tensor + mock_layer.bias.data | ||
| 304 | + self.assertTrue(torch.allclose(output, expected_output)) | ||
| 305 | + | ||
| 306 | + | ||
| 307 | + def test_apply_with_residual_bias_addition(self, mock_npu_add_rms_norm): | ||
| 308 | + """Test that bias is added correctly when residual is provided.""" | ||
| 309 | + mock_layer = MagicMock(spec=torch.nn.Module) | ||
| 310 | + mock_layer.weight = MagicMock() | ||
| 311 | + mock_layer.weight.data = torch.ones(self.hidden_size, dtype=torch.float32) | ||
| 312 | + mock_layer.bias = MagicMock() | ||
| 313 | + # Set bias to non-zero values | ||
| 314 | + mock_layer.bias.data = torch.ones(self.hidden_size, dtype=torch.float32) * 0.5 | ||
| 315 | + mock_layer.variance_epsilon = 1e-6 | ||
| 316 | + | ||
| 317 | + x = torch.randn(2, 3, self.hidden_size, dtype=torch.float32) | ||
| 318 | + residual = torch.randn(2, 3, self.hidden_size, dtype=torch.float32) | ||
| 319 | + normalized_tensor = torch.randn(2, 3, self.hidden_size, dtype=torch.float32) | ||
| 320 | + variance = torch.randn(2, 3, dtype=torch.float32) | ||
| 321 | + updated_residual = torch.randn(2, 3, self.hidden_size, dtype=torch.float32) | ||
| 322 | + mock_npu_add_rms_norm.return_value = (normalized_tensor, variance, updated_residual) | ||
| 323 | + | ||
| 324 | + output, output_residual = self.quant_method.apply(layer=mock_layer, x=x, residual=residual) | ||
| 325 | + | ||
| 326 | + # Verify bias was added | ||
| 327 | + expected_output = normalized_tensor + mock_layer.bias.data | ||
| 328 | + self.assertTrue(torch.allclose(output, expected_output)) | ||
| 329 | + | ||
| 330 | + | ||
| 331 | + def test_apply_without_residual_variance_epsilon(self, mock_npu_rms_norm): | ||
| 332 | + """Test that variance_epsilon is passed correctly when no residual.""" | ||
| 333 | + mock_layer = MagicMock(spec=torch.nn.Module) | ||
| 334 | + mock_layer.weight = MagicMock() | ||
| 335 | + mock_layer.weight.data = torch.ones(self.hidden_size, dtype=torch.float32) | ||
| 336 | + mock_layer.bias = MagicMock() | ||
| 337 | + mock_layer.bias.data = torch.zeros(self.hidden_size, dtype=torch.float32) | ||
| 338 | + mock_layer.variance_epsilon = 1e-5 # Custom epsilon | ||
| 339 | + | ||
| 340 | + x = torch.randn(2, 3, self.hidden_size, dtype=torch.float32) | ||
| 341 | + normalized_tensor = torch.randn(2, 3, self.hidden_size, dtype=torch.float32) | ||
| 342 | + variance = torch.randn(2, 3, dtype=torch.float32) | ||
| 343 | + mock_npu_rms_norm.return_value = (normalized_tensor, variance) | ||
| 344 | + | ||
| 345 | + self.quant_method.apply(layer=mock_layer, x=x) | ||
| 346 | + | ||
| 347 | + # Verify variance_epsilon was passed correctly | ||
| 348 | + call_args = mock_npu_rms_norm.call_args | ||
| 349 | + self.assertEqual(call_args[0][2], 1e-5) | ||
| 350 | + | ||
| 351 | + | ||
| 352 | + def test_apply_with_residual_variance_epsilon(self, mock_npu_add_rms_norm): | ||
| 353 | + """Test that variance_epsilon is passed correctly when residual is provided.""" | ||
| 354 | + mock_layer = MagicMock(spec=torch.nn.Module) | ||
| 355 | + mock_layer.weight = MagicMock() | ||
| 356 | + mock_layer.weight.data = torch.ones(self.hidden_size, dtype=torch.float32) | ||
| 357 | + mock_layer.bias = MagicMock() | ||
| 358 | + mock_layer.bias.data = torch.zeros(self.hidden_size, dtype=torch.float32) | ||
| 359 | + mock_layer.variance_epsilon = 1e-5 # Custom epsilon | ||
| 360 | + | ||
| 361 | + x = torch.randn(2, 3, self.hidden_size, dtype=torch.float32) | ||
| 362 | + residual = torch.randn(2, 3, self.hidden_size, dtype=torch.float32) | ||
| 363 | + normalized_tensor = torch.randn(2, 3, self.hidden_size, dtype=torch.float32) | ||
| 364 | + variance = torch.randn(2, 3, dtype=torch.float32) | ||
| 365 | + updated_residual = torch.randn(2, 3, self.hidden_size, dtype=torch.float32) | ||
| 366 | + mock_npu_add_rms_norm.return_value = (normalized_tensor, variance, updated_residual) | ||
| 367 | + | ||
| 368 | + self.quant_method.apply(layer=mock_layer, x=x, residual=residual) | ||
| 369 | + | ||
| 370 | + # Verify variance_epsilon was passed correctly | ||
| 371 | + call_args = mock_npu_add_rms_norm.call_args | ||
| 372 | + self.assertEqual(call_args[0][3], 1e-5) | ||
| 373 | + | ||
| 374 | + | ||
| 375 | + def test_apply_without_residual_return_type(self, mock_npu_rms_norm): | ||
| 376 | + """Test that apply returns a single tensor when no residual.""" | ||
| 377 | + mock_layer = MagicMock(spec=torch.nn.Module) | ||
| 378 | + mock_layer.weight = MagicMock() | ||
| 379 | + mock_layer.weight.data = torch.ones(self.hidden_size, dtype=torch.float32) | ||
| 380 | + mock_layer.bias = MagicMock() | ||
| 381 | + mock_layer.bias.data = torch.zeros(self.hidden_size, dtype=torch.float32) | ||
| 382 | + mock_layer.variance_epsilon = 1e-6 | ||
| 383 | + | ||
| 384 | + x = torch.randn(2, 3, self.hidden_size, dtype=torch.float32) | ||
| 385 | + normalized_tensor = torch.randn(2, 3, self.hidden_size, dtype=torch.float32) | ||
| 386 | + variance = torch.randn(2, 3, dtype=torch.float32) | ||
| 387 | + mock_npu_rms_norm.return_value = (normalized_tensor, variance) | ||
| 388 | + | ||
| 389 | + output = self.quant_method.apply(layer=mock_layer, x=x) | ||
| 390 | + | ||
| 391 | + # Should return a single tensor, not a tuple | ||
| 392 | + self.assertIsInstance(output, torch.Tensor) | ||
| 393 | + self.assertNotIsInstance(output, tuple) | ||
| 394 | + | ||
| 395 | + | ||
| 396 | + def test_apply_with_residual_return_type(self, mock_npu_add_rms_norm): | ||
| 397 | + """Test that apply returns a tuple when residual is provided.""" | ||
| 398 | + mock_layer = MagicMock(spec=torch.nn.Module) | ||
| 399 | + mock_layer.weight = MagicMock() | ||
| 400 | + mock_layer.weight.data = torch.ones(self.hidden_size, dtype=torch.float32) | ||
| 401 | + mock_layer.bias = MagicMock() | ||
| 402 | + mock_layer.bias.data = torch.zeros(self.hidden_size, dtype=torch.float32) | ||
| 403 | + mock_layer.variance_epsilon = 1e-6 | ||
| 404 | + | ||
| 405 | + x = torch.randn(2, 3, self.hidden_size, dtype=torch.float32) | ||
| 406 | + residual = torch.randn(2, 3, self.hidden_size, dtype=torch.float32) | ||
| 407 | + normalized_tensor = torch.randn(2, 3, self.hidden_size, dtype=torch.float32) | ||
| 408 | + variance = torch.randn(2, 3, dtype=torch.float32) | ||
| 409 | + updated_residual = torch.randn(2, 3, self.hidden_size, dtype=torch.float32) | ||
| 410 | + mock_npu_add_rms_norm.return_value = (normalized_tensor, variance, updated_residual) | ||
| 411 | + | ||
| 412 | + result = self.quant_method.apply(layer=mock_layer, x=x, residual=residual) | ||
| 413 | + | ||
| 414 | + # Should return a tuple | ||
| 415 | + self.assertIsInstance(result, tuple) | ||
| 416 | + self.assertEqual(len(result), 2) | ||
| 417 | + output, output_residual = result | ||
| 418 | + self.assertIsInstance(output, torch.Tensor) | ||
| 419 | + self.assertIsInstance(output_residual, torch.Tensor) | ||
| 420 | + | ||
| 421 | + | ||
| 422 | +if __name__ == '__main__': | ||
| 423 | + unittest.main() | ||
| 424 | + | ||