import dataclasses import json import torch from data import Vocab from train_target import train_target_wants_text from model import ( ModelArgs, MusicalPositionEmbedTransformer as MusicalPositionEmbedTransformerForTraining, ) from model_onnx import ( MusicalPositionEmbedTransformer, MusicalPositionEmbedTransformerOneStep, MusicalPositionEmbedTransformerOneStepNoCross, ) from tqdm import tqdm from embedder import Embedder import fire from data_gen import DataGenerator, Dataset, Augmenter import dataset_classes import data_aug MPET_LEN = 9 def load(vocab, config, model_file): state_dict = torch.load(model_file, map_location="cpu") params = dataclasses.replace( ModelArgs.from_config(vocab, config), enable_cross_attention=bool(config["enable_cross_attention"]), enable_flash=False, ) print(params) model = MusicalPositionEmbedTransformer(vocab, params) one_step_class = ( MusicalPositionEmbedTransformerOneStep if params.enable_cross_attention else MusicalPositionEmbedTransformerOneStepNoCross ) model_one_step = one_step_class(vocab, params) model.load_state_dict(state_dict, strict=False) model_one_step.load_state_dict(state_dict, strict=False) model.eval() model_one_step.eval() return model, model_one_step def load_training_model(vocab, config, model_file): state_dict = torch.load(model_file, map_location="cpu") train_target = config["train_target"] params = dataclasses.replace( ModelArgs.from_config(vocab, config), enable_cross_attention=train_target_wants_text(train_target), enable_flash=False, ) model = MusicalPositionEmbedTransformerForTraining( vocab, params, encoder=( Embedder(load_pretrained_weights=False) if train_target_wants_text(train_target) else None ), ) model.load_state_dict(state_dict) model.eval() return model class OnnxTool: def __init__(self, config_path, model_path): self.model_path = model_path with open(config_path, "r") as f: self._config = json.load(f) def export(self, out_path, out_one_step_path): """ Export the inference-optimized model to ONNX format. """ print(f"Loading model from {self.model_path} for export") enable_cross_attention = bool(self._config["enable_cross_attention"]) vocab = Vocab.from_config(self._config) model, model_one_step = load( vocab, self._config, self.model_path, ) device = next(model.parameters()).device num_layers = model.params.n_layers dim = model.params.dim encoder_dim = model.params.cross_attention_embedding_dim nheads = model.params.n_heads head_dim = dim // nheads inference_batch_size = model.params.inference_batch_size # these don't really matter, just need something to test + trace the model with decoder_seq_len = 123 encoder_seq_len = 32 a_xks = torch.ones( decoder_seq_len, inference_batch_size, num_layers, nheads, head_dim, dtype=torch.float, device=device, ) a_xvs = torch.ones( decoder_seq_len, inference_batch_size, num_layers, nheads, head_dim, dtype=torch.float, device=device, ) c_xks = torch.ones( num_layers, inference_batch_size, nheads, head_dim, encoder_seq_len, dtype=torch.float, device=device, ) c_xvs = torch.ones( num_layers, inference_batch_size, nheads, encoder_seq_len, head_dim, dtype=torch.float, device=device, ) x = torch.ones( inference_batch_size, decoder_seq_len, MPET_LEN, dtype=torch.long, device=device, ) xb = torch.ones( inference_batch_size, 1, MPET_LEN, dtype=torch.long, device=device ) e = torch.ones( inference_batch_size, encoder_seq_len, encoder_dim, dtype=torch.float, device=device, ) v = torch.ones( inference_batch_size, encoder_seq_len, dtype=torch.bool, device=device ) # check the model actually runs first with torch.no_grad(): model(x=x, encoder_out=e, encoder_valid=v) if enable_cross_attention: model_one_step( x=xb, a_xks=a_xks, a_xvs=a_xvs, c_xks=c_xks, c_xvs=c_xvs, encoder_valid=v, ) else: model_one_step(x=xb, a_xks=a_xks, a_xvs=a_xvs) print(f"Exporting model to {out_path} and {out_one_step_path}") torch.onnx.export( model, (x, e, v), out_path, verbose=False, input_names=["x", "encoder_out", "encoder_valid"], output_names=[ "y", "a_xks", "a_xvs", *(["c_xks", "c_xvs"] if enable_cross_attention else []), "embeddings", ], dynamic_axes={"x": [0, 1], "encoder_out": [0, 1], "encoder_valid": [0, 1]}, do_constant_folding=True, opset_version=16, ) torch.onnx.export( model_one_step, ( (xb, a_xks, a_xvs, c_xks, c_xvs, v) if enable_cross_attention else (xb, a_xks, a_xvs) ), out_one_step_path, verbose=False, input_names=[ "x", "a_xks", "a_xvs", *( ["c_xks", "c_xvs", "encoder_valid"] if enable_cross_attention else [] ), ], output_names=["y", "a_xk_news", "a_xv_news"], dynamic_axes=dict( {"x": [0], "a_xks": [0, 1], "a_xvs": [0, 1]}, **( {"c_xks": [1, 4], "c_xvs": [1, 3], "encoder_valid": [0, 1]} if enable_cross_attention else {} ), ), do_constant_folding=True, opset_version=16, ) def export_calibration_data(self, out_path, batch_size=3, batches=4096): """ Export calibration data for the inference-optimized model. Data will be a list of dicts, each dict containing tgt: the target sequences, as a BxLx9 tensor encoded: the encoder output, as a BxMx512 tensor encoded_valid: the encoder valid mask, as a BxM tensor ref_y: the reference output, as a BxLxN tensor """ vocab = Vocab.from_config(self._config) model = load_training_model(vocab, self._config, self.model_path) seq_len = int(self._config["seq_len"]) seq_len_min = int(self._config["seq_len_min"]) train_target = self._config["train_target"] train_split_idx = int(self._config["train_split_idx"]) dataset = Dataset.dynamic_from_config(self._config, "dataset_class") auger = Augmenter.dynamic_from_config(vocab, self._config, "augmenter_class") data_gen = DataGenerator( dataset, vocab, split=train_split_idx, ranksize=(0, 1), batch_size=batch_size, seq_len=seq_len, seq_len_min=seq_len_min, seq_len_max=None, train_target=train_target, parallelism=1, to_device="cpu", text_tokenize=model.encoder.tokenize, example_continuous=False, pack_batch=False, augmenter=auger, ) calibs = [] for batch in tqdm( data_gen.generate( force_batches=batches, order="random", ), total=batches, ): with torch.no_grad(): tgt = batch.tgt[ :, : torch.max( torch.nonzero(batch.tgt[:, :, 0] != model.params.padding_idx) ).item(), ] encoded = model.encoder( batch.encoder_input_ids, batch.encoder_attention_mask ) ref_y = model.forward( tgt, start_pos=0, encoder_input_ids=batch.encoder_input_ids, encoder_attention_mask=batch.encoder_attention_mask, ) calibs.append( { "tgt": tgt, "encoded": encoded, "encoded_valid": batch.encoder_attention_mask.to(dtype=torch.bool), "ref_y": ref_y, } ) torch.save(calibs, out_path) if __name__ == "__main__": fire.Fire(OnnxTool)