create_compressed_retrain_model

Applicability

Product

Supported

Atlas 350 Accelerator Card

  • QAT
    • INT8 quantization: √
  • Filter-level sparsity: √
  • 2:4 structured sparsity: x

Atlas A3 training product/Atlas A3 inference product

  • QAT
    • INT8 quantization: √
  • Filter-level sparsity: √
  • 2:4 structured sparsity: √

Atlas A2 training product/Atlas A2 inference product

  • QAT
    • INT8 quantization: √
  • Filter-level sparsity: √
  • 2:4 structured sparsity: √

Atlas 200I/500 A2 inference product

  • QAT
    • INT8 quantization: √
  • Filter-level sparsity: √
  • 2:4 structured sparsity: √

Atlas inference product

  • QAT
    • INT8 quantization: √
  • Filter-level sparsity: √
  • 2:4 structured sparsity: x

Atlas training product

  • QAT
    • INT8 quantization: √
  • Filter-level sparsity: √
  • 2:4 structured sparsity: x

Note: For the products marked with x, no error is reported when the API is called, but no performance gains are obtained.

Description

Applies to static compression combination. Compresses the input model based on the specified static compression combination configuration file. That is, prunes the input model (via either filter-level sparsity or 2:4 structured sparsity), inserts quantization operators (QAT layer for activations and weights and searchN layer) into the model, generates the sparsity record file record_file (if the configuration exists), and returns the modified torch.nn.Module model.

Prototype

1
compressed_retrain_model = create_compressed_retrain_model(model, input_data, config_defination, record_file)

Parameters

Parameter

Input/Output

Description

model

Input

PyTorch model.

Data type: torch.nn.Module

input_data

Input

Input data of the model.

A tuple.

config_defination

Input

Simplified configuration file for static compression combination.

The simplified configuration file compressed.cfg is generated based on the retrain_config_pytorch.proto file. The *.proto file is stored in /amct_pytorch/proto/ under the AMCT installation directory. For details about the parameters in the *.proto file and the generated simplified configuration file compressed.cfg, see Simplified QAT Configuration File.

A string.

record_file

Input

Path (including the file name) of the sparsity and quantization factor record file to be recorded.

A string.

Returns

Sparsifies the data based on the configuration file (if configured) and inserts torch.nn.Module of the quantization-related layer (if configured).

Restrictions

The compression combination configuration file must contain at least one of the following configurations: sparsity configuration or quantization configuration.

Example

 1
 2
 3
 4
 5
 6
 7
 8
 9
10
11
12
13
14
import amct_pytorch as amct
# Build a network for static compression combination.
model = build_model()
input_data = tuple([torch.randn(input_shape)])

# Call the static compression combination API.
record_file = os.path.join(TMP, 'compressed_record.txt')
config_defination = './compressed_cfg.cfg'

compressed_retrain_model = amct.create_compressed_retrain_model(
                                model,
                                input_data,
                                config_defination,
                                record_file)

Flush file:

Static compression combination record file record_file. If the simplified configuration file contains sparsity configuration, record_file contains sparsity record information after the function is executed.