from trt_execution import TRTModel import tensorrt as trt import torch from onnx_export import load import json from data import Vocab import time model_prefix = "model_2024-04-21" config_path = f"{model_prefix}/config.json" model_path = f"{model_prefix}/model.pt" def timing_test_trt(): with open(config_path, "r") as f: config = json.load(f) vocab = Vocab.from_config(config) ref_model, ref_model_one_step = load(vocab, config, model_path) ref_model.eval() ref_model_one_step.eval() ref_model = ref_model.cuda() ref_model_one_step = ref_model_one_step.cuda() TRT_LOGGER = trt.Logger(trt.Logger.VERBOSE) runtime = trt.Runtime(TRT_LOGGER) stream = torch.cuda.Stream() model = TRTModel(runtime, f"{model_prefix}/model.trt") context = model.create_execution_context(stream=stream) model_one_step = TRTModel(runtime, f"{model_prefix}/model_one_step.trt") context_one_step = model_one_step.create_execution_context(stream=stream) eval_data = torch.load(f"{model_prefix}/calibration.pt", map_location="cuda") with torch.no_grad(), torch.cuda.stream(stream): context.set_optimization_profile_index( model.get_profile_index_for_dim_constraint( "x", 0, eval_data[0]["tgt"].shape[0] ) ) context_one_step.set_optimization_profile_index( model_one_step.get_profile_index_for_dim_constraint( "x", 0, eval_data[0]["tgt"].shape[0] ) ) for obj in eval_data[:100]: tgt = obj["tgt"] encoded = obj["encoded"] encoded_valid = obj["encoded_valid"] ref_y = obj["ref_y"] start_point = ( torch.min( torch.count_nonzero(tgt[:, :, 0] != vocab.pad.index, dim=-1) ).item() - 3 ) t0 = time.time() ref_out = ref_model(tgt, encoded, encoded_valid) ref_dt = time.time() - t0 delta = ref_out[0] - ref_y delta[tgt[:, :, 0] == 0] = 0 t0 = time.time() out = context.eval( x=tgt.type(torch.int32), encoder_out=encoded, encoder_valid=encoded_valid, ) dt = time.time() - t0 delta_out = out["y"] - ref_y delta_out[tgt[:, :, 0] == 0] = 0 # test one step ref matches full ref a_xks0 = ref_out[1].transpose(0, 1).contiguous() a_xvs0 = ref_out[2].transpose(0, 1).contiguous() c_xks = ref_out[3].clone() c_xvs = ref_out[4].clone() y0 = ref_out[0].clone() t0 = time.time() y1, a_xks1, a_xvs1 = ref_model_one_step( tgt[:, start_point - 1 : start_point], a_xks0[: start_point - 1], a_xvs0[: start_point - 1], c_xks, c_xvs, encoded_valid, ) ref_step_dt = time.time() - t0 delta_step_ref = y0[:, start_point - 1] - y1[:, 0] delta_cache_k_ref = a_xks0[start_point - 1] - a_xks1 delta_cache_v_ref = a_xvs0[start_point - 1] - a_xvs1 # test one step matches full ref out = context.eval( x=tgt[:, :start_point].type(torch.int32), encoder_out=encoded, encoder_valid=encoded_valid, ) a_xks0 = out["a_xks"].transpose(0, 1).contiguous() a_xvs0 = out["a_xvs"].transpose(0, 1).contiguous() c_xks = out.get("c_xks", None).clone() c_xvs = out.get("c_xvs", None).clone() y0 = out["y"].clone() t0 = time.time() out_step = context_one_step.eval( x=tgt[:, start_point - 1 : start_point].type(torch.int32), a_xks=a_xks0[:-1], a_xvs=a_xvs0[:-1], c_xks=c_xks, c_xvs=c_xvs, encoder_valid=encoded_valid, ) step_dt = time.time() - t0 y1 = out_step["y"] a_xks1 = out_step["a_xk_news"] a_xvs1 = out_step["a_xv_news"] delta_step = y0[:, -1] - y1[:, 0] delta_cache_k = a_xks0[-1] - a_xks1 delta_cache_v = a_xvs0[-1] - a_xvs1 print( f"{ref_out[0].abs().max().item():.2f}", f"{delta.abs().max().item():.4f}", f"({ref_dt * 1000:.2f} ms)", f"{delta_out.abs().max().item():.4f}", f"({dt * 1000:.2f} ms)", f"r[{delta_step_ref.abs().max().item():.4f}", f"{delta_cache_k_ref.abs().max().item():.4f}", f"{delta_cache_v_ref.abs().max().item():.4f}]", f"({ref_step_dt * 1000:.2f} ms)", f"t[{delta_step.abs().max().item():.4f}", f"{delta_cache_k.abs().max().item():.4f}", f"{delta_cache_v.abs().max().item():.4f}]", f"({step_dt * 1000:.2f} ms)", ) if __name__ == "__main__": timing_test_trt()