ALTICDEV's picture
Publish Fluid-1 Pico
03fba3c verified
Raw
History Blame Contribute Delete
106 kB
program(1.3)
[buildInfo = dict<string, string>({{"coremlc-component-MIL", "3600.16.1"}, {"coremlc-version", "3600.22.1"}})]
{
func main<ios18>(tensor<fp16, [1, 1025, 1, 1]> causal_mask, tensor<fp16, [3, 1024, 3]> conv_state_in, tensor<fp16, [1, 1, 1024]> hidden_in, tensor<fp16, [2, 1, 512, 1, 1024]> kv_cache_in, tensor<int32, [1]> position_ids, tensor<fp16, [1, 1, 1024, 1]> update_mask) {
fp16 attn_scale = const()[name = string("attn_scale"), val = fp16(0x1p-3)];
tensor<fp16, [64]> layers_1_self_attn_k_layernorm_weight = const()[name = string("layers_1_self_attn_k_layernorm_weight"), val = tensor<fp16, [64]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(64)))];
tensor<fp16, [64]> layers_1_self_attn_q_layernorm_weight = const()[name = string("layers_1_self_attn_q_layernorm_weight"), val = tensor<fp16, [64]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(256)))];
tensor<fp16, [1024]> layers_0_operator_norm_weight = const()[name = string("layers_0_operator_norm_weight"), val = tensor<fp16, [1024]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(448)))];
tensor<fp16, [2048, 64]> sin_cached_palettized = constexpr_lut_to_dense(indices = tensor<uint6, [2048, 64]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(2560))), lut = tensor<fp16, [64, 1, 64, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(100928))))[name = string("sin_cached_palettized")];
tensor<fp16, [2048, 64]> cos_cached_palettized = constexpr_lut_to_dense(indices = tensor<uint6, [2048, 64]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(109184))), lut = tensor<fp16, [64, 1, 64, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(207552))))[name = string("cos_cached_palettized")];
tensor<fp16, [3072, 1024, 1, 1]> layers_0_conv_in_proj_weight_palettized = constexpr_lut_to_dense(indices = tensor<uint6, [3072, 1024, 1, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(215808))), lut = tensor<fp16, [96, 1, 1, 1, 64, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(2575168))))[name = string("layers_0_conv_in_proj_weight_palettized")];
tensor<fp16, [4608, 1024, 1, 1]> layers_0_feed_forward_w1_weight_palettized = constexpr_lut_to_dense(indices = tensor<uint6, [4608, 1024, 1, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(2587520))), lut = tensor<fp16, [144, 1, 1, 1, 64, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(6126528))))[name = string("layers_0_feed_forward_w1_weight_palettized")];
tensor<fp16, [4608, 1024, 1, 1]> layers_0_feed_forward_w3_weight_palettized = constexpr_lut_to_dense(indices = tensor<uint6, [4608, 1024, 1, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(6145024))), lut = tensor<fp16, [144, 1, 1, 1, 64, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(9684032))))[name = string("layers_0_feed_forward_w3_weight_palettized")];
tensor<fp16, [1024, 4608, 1, 1]> layers_0_feed_forward_w2_weight_palettized = constexpr_lut_to_dense(indices = tensor<uint6, [1024, 4608, 1, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(9702528))), lut = tensor<fp16, [32, 1, 1, 1, 64, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(13241536))))[name = string("layers_0_feed_forward_w2_weight_palettized")];
tensor<fp16, [1024, 1024, 1, 1]> layers_1_self_attn_q_proj_weight_palettized = constexpr_lut_to_dense(indices = tensor<uint6, [1024, 1024, 1, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(13245696))), lut = tensor<fp16, [32, 1, 1, 1, 64, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(14032192))))[name = string("layers_1_self_attn_q_proj_weight_palettized")];
tensor<fp16, [512, 1024, 1, 1]> layers_1_self_attn_k_proj_weight_palettized = constexpr_lut_to_dense(indices = tensor<uint6, [512, 1024, 1, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(14036352))), lut = tensor<fp16, [16, 1, 1, 1, 64, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(14429632))))[name = string("layers_1_self_attn_k_proj_weight_palettized")];
tensor<fp16, [512, 1024, 1, 1]> layers_1_self_attn_v_proj_weight_palettized = constexpr_lut_to_dense(indices = tensor<uint6, [512, 1024, 1, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(14431744))), lut = tensor<fp16, [16, 1, 1, 1, 64, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(14825024))))[name = string("layers_1_self_attn_v_proj_weight_palettized")];
tensor<fp16, [4608, 1024, 1, 1]> layers_1_feed_forward_w1_weight_palettized = constexpr_lut_to_dense(indices = tensor<uint6, [4608, 1024, 1, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(14827136))), lut = tensor<fp16, [144, 1, 1, 1, 64, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(18366144))))[name = string("layers_1_feed_forward_w1_weight_palettized")];
tensor<fp16, [4608, 1024, 1, 1]> layers_1_feed_forward_w3_weight_palettized = constexpr_lut_to_dense(indices = tensor<uint6, [4608, 1024, 1, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(18384640))), lut = tensor<fp16, [144, 1, 1, 1, 64, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(21923648))))[name = string("layers_1_feed_forward_w3_weight_palettized")];
tensor<fp16, [1024, 4608, 1, 1]> layers_1_feed_forward_w2_weight_palettized = constexpr_lut_to_dense(indices = tensor<uint6, [1024, 4608, 1, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(21942144))), lut = tensor<fp16, [32, 1, 1, 1, 64, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(25481152))))[name = string("layers_1_feed_forward_w2_weight_palettized")];
tensor<fp16, [3072, 1024, 1, 1]> layers_2_conv_in_proj_weight_palettized = constexpr_lut_to_dense(indices = tensor<uint6, [3072, 1024, 1, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(25485312))), lut = tensor<fp16, [96, 1, 1, 1, 64, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(27844672))))[name = string("layers_2_conv_in_proj_weight_palettized")];
tensor<fp16, [4608, 1024, 1, 1]> layers_2_feed_forward_w1_weight_palettized = constexpr_lut_to_dense(indices = tensor<uint6, [4608, 1024, 1, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(27857024))), lut = tensor<fp16, [144, 1, 1, 1, 64, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(31396032))))[name = string("layers_2_feed_forward_w1_weight_palettized")];
tensor<fp16, [4608, 1024, 1, 1]> layers_2_feed_forward_w3_weight_palettized = constexpr_lut_to_dense(indices = tensor<uint6, [4608, 1024, 1, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(31414528))), lut = tensor<fp16, [144, 1, 1, 1, 64, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(34953536))))[name = string("layers_2_feed_forward_w3_weight_palettized")];
tensor<fp16, [1024, 4608, 1, 1]> layers_2_feed_forward_w2_weight_palettized = constexpr_lut_to_dense(indices = tensor<uint6, [1024, 4608, 1, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(34972032))), lut = tensor<fp16, [32, 1, 1, 1, 64, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(38511040))))[name = string("layers_2_feed_forward_w2_weight_palettized")];
tensor<fp16, [3072, 1024, 1, 1]> layers_3_conv_in_proj_weight_palettized = constexpr_lut_to_dense(indices = tensor<uint6, [3072, 1024, 1, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(38515200))), lut = tensor<fp16, [96, 1, 1, 1, 64, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(40874560))))[name = string("layers_3_conv_in_proj_weight_palettized")];
tensor<fp16, [4608, 1024, 1, 1]> layers_3_feed_forward_w1_weight_palettized = constexpr_lut_to_dense(indices = tensor<uint6, [4608, 1024, 1, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(40886912))), lut = tensor<fp16, [144, 1, 1, 1, 64, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(44425920))))[name = string("layers_3_feed_forward_w1_weight_palettized")];
tensor<fp16, [4608, 1024, 1, 1]> layers_3_feed_forward_w3_weight_palettized = constexpr_lut_to_dense(indices = tensor<uint6, [4608, 1024, 1, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(44444416))), lut = tensor<fp16, [144, 1, 1, 1, 64, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(47983424))))[name = string("layers_3_feed_forward_w3_weight_palettized")];
tensor<fp16, [1024, 4608, 1, 1]> layers_3_feed_forward_w2_weight_palettized = constexpr_lut_to_dense(indices = tensor<uint6, [1024, 4608, 1, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(48001920))), lut = tensor<fp16, [32, 1, 1, 1, 64, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(51540928))))[name = string("layers_3_feed_forward_w2_weight_palettized")];
int32 var_158_batch_dims_0 = const()[name = string("op_158_batch_dims_0"), val = int32(0)];
bool var_158_validate_indices_0 = const()[name = string("op_158_validate_indices_0"), val = bool(false)];
string position_ids_to_int16_dtype_0 = const()[name = string("position_ids_to_int16_dtype_0"), val = string("int16")];
string cast_15_dtype_0 = const()[name = string("cast_15_dtype_0"), val = string("int32")];
int32 greater_equal_0_y_0 = const()[name = string("greater_equal_0_y_0"), val = int32(0)];
tensor<int16, [1]> position_ids_to_int16 = cast(dtype = position_ids_to_int16_dtype_0, x = position_ids)[name = string("cast_5")];
tensor<int32, [1]> cast_15 = cast(dtype = cast_15_dtype_0, x = position_ids_to_int16)[name = string("cast_4")];
tensor<bool, [1]> greater_equal_0 = greater_equal(x = cast_15, y = greater_equal_0_y_0)[name = string("greater_equal_0")];
int32 slice_by_index_0 = const()[name = string("slice_by_index_0"), val = int32(2048)];
tensor<int32, [1]> add_0 = add(x = cast_15, y = slice_by_index_0)[name = string("add_0")];
tensor<int32, [1]> select_0 = select(a = cast_15, b = add_0, cond = greater_equal_0)[name = string("select_0")];
string select_0_to_int16_dtype_0 = const()[name = string("select_0_to_int16_dtype_0"), val = string("int16")];
string cast_0_dtype_0 = const()[name = string("cast_0_dtype_0"), val = string("int32")];
int32 greater_equal_0_y_0_1 = const()[name = string("greater_equal_0_y_0_1"), val = int32(0)];
tensor<int16, [1]> select_0_to_int16 = cast(dtype = select_0_to_int16_dtype_0, x = select_0)[name = string("cast_3")];
tensor<int32, [1]> cast_0 = cast(dtype = cast_0_dtype_0, x = select_0_to_int16)[name = string("cast_2")];
tensor<bool, [1]> greater_equal_0_1 = greater_equal(x = cast_0, y = greater_equal_0_y_0_1)[name = string("greater_equal_0_1")];
int32 slice_by_index_0_1 = const()[name = string("slice_by_index_0_1"), val = int32(2048)];
tensor<int32, [1]> add_0_1 = add(x = cast_0, y = slice_by_index_0_1)[name = string("add_0_1")];
tensor<int32, [1]> select_0_1 = select(a = cast_0, b = add_0_1, cond = greater_equal_0_1)[name = string("select_0_1")];
int32 op_158_cast_uint16_cast_uint16_axis_0 = const()[name = string("op_158_cast_uint16_cast_uint16_axis_0"), val = int32(0)];
tensor<fp16, [1, 64]> op_158_cast_uint16_cast_uint16 = gather(axis = op_158_cast_uint16_cast_uint16_axis_0, batch_dims = var_158_batch_dims_0, indices = select_0_1, validate_indices = var_158_validate_indices_0, x = cos_cached_palettized)[name = string("op_158_cast_uint16_cast_uint16")];
tensor<int32, [4]> var_163 = const()[name = string("op_163"), val = tensor<int32, [4]>([1, 1, 1, 64])];
tensor<fp16, [1, 1, 1, 64]> cos = reshape(shape = var_163, x = op_158_cast_uint16_cast_uint16)[name = string("cos")];
int32 var_165 = const()[name = string("op_165"), val = int32(0)];
int32 var_166_batch_dims_0 = const()[name = string("op_166_batch_dims_0"), val = int32(0)];
bool var_166_validate_indices_0 = const()[name = string("op_166_validate_indices_0"), val = bool(false)];
string position_ids_to_uint16_dtype_0 = const()[name = string("position_ids_to_uint16_dtype_0"), val = string("uint16")];
tensor<uint16, [1]> position_ids_to_uint16 = cast(dtype = position_ids_to_uint16_dtype_0, x = position_ids)[name = string("cast_1")];
tensor<fp16, [1, 64]> var_166_cast_uint16 = gather(axis = var_165, batch_dims = var_166_batch_dims_0, indices = position_ids_to_uint16, validate_indices = var_166_validate_indices_0, x = sin_cached_palettized)[name = string("op_166_cast_uint16")];
tensor<int32, [4]> var_171 = const()[name = string("op_171"), val = tensor<int32, [4]>([1, 1, 1, 64])];
tensor<fp16, [1, 1, 1, 64]> sin = reshape(shape = var_171, x = var_166_cast_uint16)[name = string("sin")];
fp16 const_0_promoted = const()[name = string("const_0_promoted"), val = fp16(-0x1p+0)];
tensor<fp16, [1, 1, 1024]> var_173 = mul(x = hidden_in, y = const_0_promoted)[name = string("op_173")];
int32 var_175 = const()[name = string("op_175"), val = int32(-1)];
bool input_1_interleave_0 = const()[name = string("input_1_interleave_0"), val = bool(false)];
tensor<fp16, [1, 1, 2048]> input_1 = concat(axis = var_175, interleave = input_1_interleave_0, values = (hidden_in, var_173))[name = string("input_1")];
tensor<int32, [1]> normed_1_axes_0 = const()[name = string("normed_1_axes_0"), val = tensor<int32, [1]>([-1])];
fp16 var_181_to_fp16 = const()[name = string("op_181_to_fp16"), val = fp16(0x1.5p-17)];
tensor<fp16, [1, 1, 2048]> normed_1_cast_fp16 = layer_norm(axes = normed_1_axes_0, epsilon = var_181_to_fp16, x = input_1)[name = string("normed_1_cast_fp16")];
tensor<int32, [2]> var_184_split_sizes_0 = const()[name = string("op_184_split_sizes_0"), val = tensor<int32, [2]>([1024, 1024])];
int32 var_184_axis_0 = const()[name = string("op_184_axis_0"), val = int32(-1)];
tensor<fp16, [1, 1, 1024]> var_184_0, tensor<fp16, [1, 1, 1024]> var_184_1 = split(axis = var_184_axis_0, split_sizes = var_184_split_sizes_0, x = normed_1_cast_fp16)[name = string("op_184")];
tensor<fp16, [1, 1, 1024]> hidden_states_1 = mul(x = var_184_0, y = layers_0_operator_norm_weight)[name = string("hidden_states_1")];
tensor<int32, [3]> var_190 = const()[name = string("op_190"), val = tensor<int32, [3]>([0, 2, 1])];
tensor<int32, [1]> var_193_axes_0 = const()[name = string("op_193_axes_0"), val = tensor<int32, [1]>([2])];
tensor<fp16, [1, 1024, 1]> var_191 = transpose(perm = var_190, x = hidden_states_1)[name = string("transpose_20")];
tensor<fp16, [1, 1024, 1, 1]> var_193 = expand_dims(axes = var_193_axes_0, x = var_191)[name = string("op_193")];
string BCx_1_pad_type_0 = const()[name = string("BCx_1_pad_type_0"), val = string("valid")];
tensor<int32, [2]> BCx_1_strides_0 = const()[name = string("BCx_1_strides_0"), val = tensor<int32, [2]>([1, 1])];
tensor<int32, [4]> BCx_1_pad_0 = const()[name = string("BCx_1_pad_0"), val = tensor<int32, [4]>([0, 0, 0, 0])];
tensor<int32, [2]> BCx_1_dilations_0 = const()[name = string("BCx_1_dilations_0"), val = tensor<int32, [2]>([1, 1])];
int32 BCx_1_groups_0 = const()[name = string("BCx_1_groups_0"), val = int32(1)];
tensor<fp16, [1, 3072, 1, 1]> BCx_1 = conv(dilations = BCx_1_dilations_0, groups = BCx_1_groups_0, pad = BCx_1_pad_0, pad_type = BCx_1_pad_type_0, strides = BCx_1_strides_0, weight = layers_0_conv_in_proj_weight_palettized, x = var_193)[name = string("BCx_1")];
tensor<int32, [3]> var_210_split_sizes_0 = const()[name = string("op_210_split_sizes_0"), val = tensor<int32, [3]>([1024, 1024, 1024])];
int32 var_210_axis_0 = const()[name = string("op_210_axis_0"), val = int32(1)];
tensor<fp16, [1, 1024, 1, 1]> var_210_0, tensor<fp16, [1, 1024, 1, 1]> var_210_1, tensor<fp16, [1, 1024, 1, 1]> var_210_2 = split(axis = var_210_axis_0, split_sizes = var_210_split_sizes_0, x = BCx_1)[name = string("op_210")];
tensor<fp16, [1, 1024, 1, 1]> Bx_1 = mul(x = var_210_0, y = var_210_2)[name = string("Bx_1")];
tensor<int32, [3]> var_216_begin_0 = const()[name = string("op_216_begin_0"), val = tensor<int32, [3]>([0, 0, 0])];
tensor<int32, [3]> var_216_end_0 = const()[name = string("op_216_end_0"), val = tensor<int32, [3]>([1, 1024, 3])];
tensor<bool, [3]> var_216_end_mask_0 = const()[name = string("op_216_end_mask_0"), val = tensor<bool, [3]>([false, true, true])];
tensor<bool, [3]> var_216_squeeze_mask_0 = const()[name = string("op_216_squeeze_mask_0"), val = tensor<bool, [3]>([true, false, false])];
tensor<fp16, [1024, 3]> var_216_cast_fp16 = slice_by_index(begin = var_216_begin_0, end = var_216_end_0, end_mask = var_216_end_mask_0, squeeze_mask = var_216_squeeze_mask_0, x = conv_state_in)[name = string("op_216_cast_fp16")];
tensor<int32, [1]> var_218_axes_0 = const()[name = string("op_218_axes_0"), val = tensor<int32, [1]>([0])];
tensor<fp16, [1, 1024, 3]> var_218_cast_fp16 = expand_dims(axes = var_218_axes_0, x = var_216_cast_fp16)[name = string("op_218_cast_fp16")];
tensor<int32, [1]> slot_1_axes_0 = const()[name = string("slot_1_axes_0"), val = tensor<int32, [1]>([2])];
tensor<fp16, [1, 1024, 1, 3]> slot_1_cast_fp16 = expand_dims(axes = slot_1_axes_0, x = var_218_cast_fp16)[name = string("slot_1_cast_fp16")];
tensor<int32, [4]> live_tail_1_begin_0 = const()[name = string("live_tail_1_begin_0"), val = tensor<int32, [4]>([0, 0, 0, 1])];
tensor<int32, [4]> live_tail_1_end_0 = const()[name = string("live_tail_1_end_0"), val = tensor<int32, [4]>([1, 1024, 1, 1])];
tensor<bool, [4]> live_tail_1_end_mask_0 = const()[name = string("live_tail_1_end_mask_0"), val = tensor<bool, [4]>([true, true, true, true])];
tensor<fp16, [1, 1024, 1, 2]> live_tail_1_cast_fp16 = slice_by_index(begin = live_tail_1_begin_0, end = live_tail_1_end_0, end_mask = live_tail_1_end_mask_0, x = slot_1_cast_fp16)[name = string("live_tail_1_cast_fp16")];
int32 var_227 = const()[name = string("op_227"), val = int32(-1)];
bool new_state_1_interleave_0 = const()[name = string("new_state_1_interleave_0"), val = bool(false)];
tensor<fp16, [1, 1024, 1, 3]> new_state_1_cast_fp16 = concat(axis = var_227, interleave = new_state_1_interleave_0, values = (live_tail_1_cast_fp16, Bx_1))[name = string("new_state_1_cast_fp16")];
tensor<int32, [1]> var_230_axes_0 = const()[name = string("op_230_axes_0"), val = tensor<int32, [1]>([0])];
tensor<fp16, [1024, 1, 3]> var_230_cast_fp16 = squeeze(axes = var_230_axes_0, x = new_state_1_cast_fp16)[name = string("op_230_cast_fp16")];
tensor<int32, [1]> var_232_axes_0 = const()[name = string("op_232_axes_0"), val = tensor<int32, [1]>([1])];
tensor<fp16, [1024, 3]> var_232_cast_fp16 = squeeze(axes = var_232_axes_0, x = var_230_cast_fp16)[name = string("op_232_cast_fp16")];
string conv_out_1_pad_type_0 = const()[name = string("conv_out_1_pad_type_0"), val = string("valid")];
int32 conv_out_1_groups_0 = const()[name = string("conv_out_1_groups_0"), val = int32(1024)];
tensor<int32, [2]> conv_out_1_strides_0 = const()[name = string("conv_out_1_strides_0"), val = tensor<int32, [2]>([1, 1])];
tensor<int32, [4]> conv_out_1_pad_0 = const()[name = string("conv_out_1_pad_0"), val = tensor<int32, [4]>([0, 0, 0, 0])];
tensor<int32, [2]> conv_out_1_dilations_0 = const()[name = string("conv_out_1_dilations_0"), val = tensor<int32, [2]>([1, 1])];
tensor<fp16, [1024, 1, 1, 3]> layers_0_conv_conv_weight_promoted_to_fp16_palettized = constexpr_lut_to_dense(indices = tensor<uint6, [1024, 1, 1, 3]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(51545088))), lut = tensor<fp16, [32, 1, 1, 1, 64, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(51547456))))[name = string("layers_0_conv_conv_weight_promoted_to_fp16_palettized")];
tensor<fp16, [1, 1024, 1, 1]> conv_out_1_cast_fp16 = conv(dilations = conv_out_1_dilations_0, groups = conv_out_1_groups_0, pad = conv_out_1_pad_0, pad_type = conv_out_1_pad_type_0, strides = conv_out_1_strides_0, weight = layers_0_conv_conv_weight_promoted_to_fp16_palettized, x = new_state_1_cast_fp16)[name = string("conv_out_1_cast_fp16")];
tensor<fp16, [1, 1024, 1, 1]> input_5_cast_fp16 = mul(x = var_210_1, y = conv_out_1_cast_fp16)[name = string("input_5_cast_fp16")];
string y_1_pad_type_0 = const()[name = string("y_1_pad_type_0"), val = string("valid")];
tensor<int32, [2]> y_1_strides_0 = const()[name = string("y_1_strides_0"), val = tensor<int32, [2]>([1, 1])];
tensor<int32, [4]> y_1_pad_0 = const()[name = string("y_1_pad_0"), val = tensor<int32, [4]>([0, 0, 0, 0])];
tensor<int32, [2]> y_1_dilations_0 = const()[name = string("y_1_dilations_0"), val = tensor<int32, [2]>([1, 1])];
int32 y_1_groups_0 = const()[name = string("y_1_groups_0"), val = int32(1)];
tensor<fp16, [1024, 1024, 1, 1]> layers_0_conv_out_proj_weight_promoted_to_fp16_palettized = constexpr_lut_to_dense(indices = tensor<uint6, [1024, 1024, 1, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(51551616))), lut = tensor<fp16, [32, 1, 1, 1, 64, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(52338112))))[name = string("layers_0_conv_out_proj_weight_promoted_to_fp16_palettized")];
tensor<fp16, [1, 1024, 1, 1]> y_1_cast_fp16 = conv(dilations = y_1_dilations_0, groups = y_1_groups_0, pad = y_1_pad_0, pad_type = y_1_pad_type_0, strides = y_1_strides_0, weight = layers_0_conv_out_proj_weight_promoted_to_fp16_palettized, x = input_5_cast_fp16)[name = string("y_1_cast_fp16")];
tensor<int32, [1]> var_258_axes_0 = const()[name = string("op_258_axes_0"), val = tensor<int32, [1]>([2])];
tensor<fp16, [1, 1024, 1]> var_258_cast_fp16 = squeeze(axes = var_258_axes_0, x = y_1_cast_fp16)[name = string("op_258_cast_fp16")];
tensor<int32, [3]> var_262 = const()[name = string("op_262"), val = tensor<int32, [3]>([0, 2, 1])];
tensor<fp16, [1, 1, 1024]> op_out_1_cast_fp16 = transpose(perm = var_262, x = var_258_cast_fp16)[name = string("transpose_19")];
tensor<fp16, [1, 1, 1024]> x_3_cast_fp16 = add(x = hidden_in, y = op_out_1_cast_fp16)[name = string("x_3_cast_fp16")];
fp16 const_1_promoted_to_fp16 = const()[name = string("const_1_promoted_to_fp16"), val = fp16(-0x1p+0)];
tensor<fp16, [1, 1, 1024]> var_266_cast_fp16 = mul(x = x_3_cast_fp16, y = const_1_promoted_to_fp16)[name = string("op_266_cast_fp16")];
int32 var_268 = const()[name = string("op_268"), val = int32(-1)];
bool input_7_interleave_0 = const()[name = string("input_7_interleave_0"), val = bool(false)];
tensor<fp16, [1, 1, 2048]> input_7_cast_fp16 = concat(axis = var_268, interleave = input_7_interleave_0, values = (x_3_cast_fp16, var_266_cast_fp16))[name = string("input_7_cast_fp16")];
tensor<int32, [1]> normed_3_axes_0 = const()[name = string("normed_3_axes_0"), val = tensor<int32, [1]>([-1])];
fp16 var_274_to_fp16 = const()[name = string("op_274_to_fp16"), val = fp16(0x1.5p-17)];
tensor<fp16, [1, 1, 2048]> normed_3_cast_fp16 = layer_norm(axes = normed_3_axes_0, epsilon = var_274_to_fp16, x = input_7_cast_fp16)[name = string("normed_3_cast_fp16")];
tensor<int32, [2]> var_277_split_sizes_0 = const()[name = string("op_277_split_sizes_0"), val = tensor<int32, [2]>([1024, 1024])];
int32 var_277_axis_0 = const()[name = string("op_277_axis_0"), val = int32(-1)];
tensor<fp16, [1, 1, 1024]> var_277_cast_fp16_0, tensor<fp16, [1, 1, 1024]> var_277_cast_fp16_1 = split(axis = var_277_axis_0, split_sizes = var_277_split_sizes_0, x = normed_3_cast_fp16)[name = string("op_277_cast_fp16")];
tensor<fp16, [1024]> layers_0_ffn_norm_weight_promoted_to_fp16 = const()[name = string("layers_0_ffn_norm_weight_promoted_to_fp16"), val = tensor<fp16, [1024]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(52342272)))];
tensor<fp16, [1, 1, 1024]> normed_5_cast_fp16 = mul(x = var_277_cast_fp16_0, y = layers_0_ffn_norm_weight_promoted_to_fp16)[name = string("normed_5_cast_fp16")];
tensor<int32, [3]> var_283 = const()[name = string("op_283"), val = tensor<int32, [3]>([0, 2, 1])];
tensor<int32, [1]> var_286_axes_0 = const()[name = string("op_286_axes_0"), val = tensor<int32, [1]>([2])];
tensor<fp16, [1, 1024, 1]> var_284_cast_fp16 = transpose(perm = var_283, x = normed_5_cast_fp16)[name = string("transpose_18")];
tensor<fp16, [1, 1024, 1, 1]> var_286_cast_fp16 = expand_dims(axes = var_286_axes_0, x = var_284_cast_fp16)[name = string("op_286_cast_fp16")];
string input_11_pad_type_0 = const()[name = string("input_11_pad_type_0"), val = string("valid")];
tensor<int32, [2]> input_11_strides_0 = const()[name = string("input_11_strides_0"), val = tensor<int32, [2]>([1, 1])];
tensor<int32, [4]> input_11_pad_0 = const()[name = string("input_11_pad_0"), val = tensor<int32, [4]>([0, 0, 0, 0])];
tensor<int32, [2]> input_11_dilations_0 = const()[name = string("input_11_dilations_0"), val = tensor<int32, [2]>([1, 1])];
int32 input_11_groups_0 = const()[name = string("input_11_groups_0"), val = int32(1)];
tensor<fp16, [1, 4608, 1, 1]> input_11 = conv(dilations = input_11_dilations_0, groups = input_11_groups_0, pad = input_11_pad_0, pad_type = input_11_pad_type_0, strides = input_11_strides_0, weight = layers_0_feed_forward_w1_weight_palettized, x = var_286_cast_fp16)[name = string("input_11")];
string b_1_pad_type_0 = const()[name = string("b_1_pad_type_0"), val = string("valid")];
tensor<int32, [2]> b_1_strides_0 = const()[name = string("b_1_strides_0"), val = tensor<int32, [2]>([1, 1])];
tensor<int32, [4]> b_1_pad_0 = const()[name = string("b_1_pad_0"), val = tensor<int32, [4]>([0, 0, 0, 0])];
tensor<int32, [2]> b_1_dilations_0 = const()[name = string("b_1_dilations_0"), val = tensor<int32, [2]>([1, 1])];
int32 b_1_groups_0 = const()[name = string("b_1_groups_0"), val = int32(1)];
tensor<fp16, [1, 4608, 1, 1]> b_1 = conv(dilations = b_1_dilations_0, groups = b_1_groups_0, pad = b_1_pad_0, pad_type = b_1_pad_type_0, strides = b_1_strides_0, weight = layers_0_feed_forward_w3_weight_palettized, x = var_286_cast_fp16)[name = string("b_1")];
tensor<fp16, [1, 4608, 1, 1]> var_314 = silu(x = input_11)[name = string("op_314")];
tensor<fp16, [1, 4608, 1, 1]> input_13 = mul(x = var_314, y = b_1)[name = string("input_13")];
string mlp_1_pad_type_0 = const()[name = string("mlp_1_pad_type_0"), val = string("valid")];
tensor<int32, [2]> mlp_1_strides_0 = const()[name = string("mlp_1_strides_0"), val = tensor<int32, [2]>([1, 1])];
tensor<int32, [4]> mlp_1_pad_0 = const()[name = string("mlp_1_pad_0"), val = tensor<int32, [4]>([0, 0, 0, 0])];
tensor<int32, [2]> mlp_1_dilations_0 = const()[name = string("mlp_1_dilations_0"), val = tensor<int32, [2]>([1, 1])];
int32 mlp_1_groups_0 = const()[name = string("mlp_1_groups_0"), val = int32(1)];
tensor<fp16, [1, 1024, 1, 1]> mlp_1 = conv(dilations = mlp_1_dilations_0, groups = mlp_1_groups_0, pad = mlp_1_pad_0, pad_type = mlp_1_pad_type_0, strides = mlp_1_strides_0, weight = layers_0_feed_forward_w2_weight_palettized, x = input_13)[name = string("mlp_1")];
tensor<int32, [1]> var_328_axes_0 = const()[name = string("op_328_axes_0"), val = tensor<int32, [1]>([2])];
tensor<fp16, [1, 1024, 1]> var_328 = squeeze(axes = var_328_axes_0, x = mlp_1)[name = string("op_328")];
tensor<int32, [3]> var_332 = const()[name = string("op_332"), val = tensor<int32, [3]>([0, 2, 1])];
tensor<fp16, [1, 1, 1024]> mlp_3 = transpose(perm = var_332, x = var_328)[name = string("transpose_17")];
tensor<fp16, [1, 1, 1024]> x_5_cast_fp16 = add(x = x_3_cast_fp16, y = mlp_3)[name = string("x_5_cast_fp16")];
fp16 const_2_promoted_to_fp16 = const()[name = string("const_2_promoted_to_fp16"), val = fp16(-0x1p+0)];
tensor<fp16, [1, 1, 1024]> var_336_cast_fp16 = mul(x = x_5_cast_fp16, y = const_2_promoted_to_fp16)[name = string("op_336_cast_fp16")];
int32 var_338 = const()[name = string("op_338"), val = int32(-1)];
bool input_15_interleave_0 = const()[name = string("input_15_interleave_0"), val = bool(false)];
tensor<fp16, [1, 1, 2048]> input_15_cast_fp16 = concat(axis = var_338, interleave = input_15_interleave_0, values = (x_5_cast_fp16, var_336_cast_fp16))[name = string("input_15_cast_fp16")];
tensor<int32, [1]> normed_7_axes_0 = const()[name = string("normed_7_axes_0"), val = tensor<int32, [1]>([-1])];
fp16 var_344_to_fp16 = const()[name = string("op_344_to_fp16"), val = fp16(0x1.5p-17)];
tensor<fp16, [1, 1, 2048]> normed_7_cast_fp16 = layer_norm(axes = normed_7_axes_0, epsilon = var_344_to_fp16, x = input_15_cast_fp16)[name = string("normed_7_cast_fp16")];
tensor<int32, [2]> var_347_split_sizes_0 = const()[name = string("op_347_split_sizes_0"), val = tensor<int32, [2]>([1024, 1024])];
int32 var_347_axis_0 = const()[name = string("op_347_axis_0"), val = int32(-1)];
tensor<fp16, [1, 1, 1024]> var_347_cast_fp16_0, tensor<fp16, [1, 1, 1024]> var_347_cast_fp16_1 = split(axis = var_347_axis_0, split_sizes = var_347_split_sizes_0, x = normed_7_cast_fp16)[name = string("op_347_cast_fp16")];
tensor<fp16, [1024]> layers_1_operator_norm_weight_promoted_to_fp16 = const()[name = string("layers_1_operator_norm_weight_promoted_to_fp16"), val = tensor<fp16, [1024]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(52344384)))];
tensor<fp16, [1, 1, 1024]> hidden_states_3_cast_fp16 = mul(x = var_347_cast_fp16_0, y = layers_1_operator_norm_weight_promoted_to_fp16)[name = string("hidden_states_3_cast_fp16")];
tensor<int32, [3]> var_353 = const()[name = string("op_353"), val = tensor<int32, [3]>([0, 2, 1])];
tensor<int32, [1]> var_356_axes_0 = const()[name = string("op_356_axes_0"), val = tensor<int32, [1]>([2])];
tensor<fp16, [1, 1024, 1]> var_354_cast_fp16 = transpose(perm = var_353, x = hidden_states_3_cast_fp16)[name = string("transpose_16")];
tensor<fp16, [1, 1024, 1, 1]> var_356_cast_fp16 = expand_dims(axes = var_356_axes_0, x = var_354_cast_fp16)[name = string("op_356_cast_fp16")];
string var_372_pad_type_0 = const()[name = string("op_372_pad_type_0"), val = string("valid")];
tensor<int32, [2]> var_372_strides_0 = const()[name = string("op_372_strides_0"), val = tensor<int32, [2]>([1, 1])];
tensor<int32, [4]> var_372_pad_0 = const()[name = string("op_372_pad_0"), val = tensor<int32, [4]>([0, 0, 0, 0])];
tensor<int32, [2]> var_372_dilations_0 = const()[name = string("op_372_dilations_0"), val = tensor<int32, [2]>([1, 1])];
int32 var_372_groups_0 = const()[name = string("op_372_groups_0"), val = int32(1)];
tensor<fp16, [1, 1024, 1, 1]> var_372 = conv(dilations = var_372_dilations_0, groups = var_372_groups_0, pad = var_372_pad_0, pad_type = var_372_pad_type_0, strides = var_372_strides_0, weight = layers_1_self_attn_q_proj_weight_palettized, x = var_356_cast_fp16)[name = string("op_372")];
tensor<int32, [4]> var_377 = const()[name = string("op_377"), val = tensor<int32, [4]>([1, 16, 64, 1])];
tensor<fp16, [1, 16, 64, 1]> var_378 = reshape(shape = var_377, x = var_372)[name = string("op_378")];
tensor<int32, [4]> var_383 = const()[name = string("op_383"), val = tensor<int32, [4]>([0, 1, 3, 2])];
string var_400_pad_type_0 = const()[name = string("op_400_pad_type_0"), val = string("valid")];
tensor<int32, [2]> var_400_strides_0 = const()[name = string("op_400_strides_0"), val = tensor<int32, [2]>([1, 1])];
tensor<int32, [4]> var_400_pad_0 = const()[name = string("op_400_pad_0"), val = tensor<int32, [4]>([0, 0, 0, 0])];
tensor<int32, [2]> var_400_dilations_0 = const()[name = string("op_400_dilations_0"), val = tensor<int32, [2]>([1, 1])];
int32 var_400_groups_0 = const()[name = string("op_400_groups_0"), val = int32(1)];
tensor<fp16, [1, 512, 1, 1]> var_400 = conv(dilations = var_400_dilations_0, groups = var_400_groups_0, pad = var_400_pad_0, pad_type = var_400_pad_type_0, strides = var_400_strides_0, weight = layers_1_self_attn_k_proj_weight_palettized, x = var_356_cast_fp16)[name = string("op_400")];
tensor<int32, [4]> var_405 = const()[name = string("op_405"), val = tensor<int32, [4]>([1, 8, 64, 1])];
tensor<fp16, [1, 8, 64, 1]> var_406 = reshape(shape = var_405, x = var_400)[name = string("op_406")];
tensor<int32, [4]> var_411 = const()[name = string("op_411"), val = tensor<int32, [4]>([0, 1, 3, 2])];
string var_428_pad_type_0 = const()[name = string("op_428_pad_type_0"), val = string("valid")];
tensor<int32, [2]> var_428_strides_0 = const()[name = string("op_428_strides_0"), val = tensor<int32, [2]>([1, 1])];
tensor<int32, [4]> var_428_pad_0 = const()[name = string("op_428_pad_0"), val = tensor<int32, [4]>([0, 0, 0, 0])];
tensor<int32, [2]> var_428_dilations_0 = const()[name = string("op_428_dilations_0"), val = tensor<int32, [2]>([1, 1])];
int32 var_428_groups_0 = const()[name = string("op_428_groups_0"), val = int32(1)];
tensor<fp16, [1, 512, 1, 1]> var_428 = conv(dilations = var_428_dilations_0, groups = var_428_groups_0, pad = var_428_pad_0, pad_type = var_428_pad_type_0, strides = var_428_strides_0, weight = layers_1_self_attn_v_proj_weight_palettized, x = var_356_cast_fp16)[name = string("op_428")];
fp16 const_3_promoted = const()[name = string("const_3_promoted"), val = fp16(-0x1p+0)];
tensor<fp16, [1, 16, 1, 64]> var_384 = transpose(perm = var_383, x = var_378)[name = string("transpose_15")];
tensor<fp16, [1, 16, 1, 64]> var_446 = mul(x = var_384, y = const_3_promoted)[name = string("op_446")];
int32 var_448 = const()[name = string("op_448"), val = int32(-1)];
bool input_19_interleave_0 = const()[name = string("input_19_interleave_0"), val = bool(false)];
tensor<fp16, [1, 16, 1, 128]> input_19 = concat(axis = var_448, interleave = input_19_interleave_0, values = (var_384, var_446))[name = string("input_19")];
tensor<int32, [1]> normed_9_axes_0 = const()[name = string("normed_9_axes_0"), val = tensor<int32, [1]>([-1])];
fp16 var_454_to_fp16 = const()[name = string("op_454_to_fp16"), val = fp16(0x1.5p-17)];
tensor<fp16, [1, 16, 1, 128]> normed_9_cast_fp16 = layer_norm(axes = normed_9_axes_0, epsilon = var_454_to_fp16, x = input_19)[name = string("normed_9_cast_fp16")];
tensor<int32, [2]> var_457_split_sizes_0 = const()[name = string("op_457_split_sizes_0"), val = tensor<int32, [2]>([64, 64])];
int32 var_457_axis_0 = const()[name = string("op_457_axis_0"), val = int32(-1)];
tensor<fp16, [1, 16, 1, 64]> var_457_0, tensor<fp16, [1, 16, 1, 64]> var_457_1 = split(axis = var_457_axis_0, split_sizes = var_457_split_sizes_0, x = normed_9_cast_fp16)[name = string("op_457")];
tensor<fp16, [1, 16, 1, 64]> q_1 = mul(x = var_457_0, y = layers_1_self_attn_q_layernorm_weight)[name = string("q_1")];
fp16 const_4_promoted = const()[name = string("const_4_promoted"), val = fp16(-0x1p+0)];
tensor<fp16, [1, 8, 1, 64]> var_412 = transpose(perm = var_411, x = var_406)[name = string("transpose_14")];
tensor<fp16, [1, 8, 1, 64]> var_460 = mul(x = var_412, y = const_4_promoted)[name = string("op_460")];
int32 var_462 = const()[name = string("op_462"), val = int32(-1)];
bool input_21_interleave_0 = const()[name = string("input_21_interleave_0"), val = bool(false)];
tensor<fp16, [1, 8, 1, 128]> input_21 = concat(axis = var_462, interleave = input_21_interleave_0, values = (var_412, var_460))[name = string("input_21")];
tensor<int32, [1]> normed_11_axes_0 = const()[name = string("normed_11_axes_0"), val = tensor<int32, [1]>([-1])];
fp16 var_468_to_fp16 = const()[name = string("op_468_to_fp16"), val = fp16(0x1.5p-17)];
tensor<fp16, [1, 8, 1, 128]> normed_11_cast_fp16 = layer_norm(axes = normed_11_axes_0, epsilon = var_468_to_fp16, x = input_21)[name = string("normed_11_cast_fp16")];
tensor<int32, [2]> var_471_split_sizes_0 = const()[name = string("op_471_split_sizes_0"), val = tensor<int32, [2]>([64, 64])];
int32 var_471_axis_0 = const()[name = string("op_471_axis_0"), val = int32(-1)];
tensor<fp16, [1, 8, 1, 64]> var_471_0, tensor<fp16, [1, 8, 1, 64]> var_471_1 = split(axis = var_471_axis_0, split_sizes = var_471_split_sizes_0, x = normed_11_cast_fp16)[name = string("op_471")];
tensor<fp16, [1, 8, 1, 64]> k_1 = mul(x = var_471_0, y = layers_1_self_attn_k_layernorm_weight)[name = string("k_1")];
tensor<fp16, [1, 16, 1, 64]> var_474 = mul(x = q_1, y = cos)[name = string("op_474")];
tensor<int32, [2]> var_475_split_sizes_0 = const()[name = string("op_475_split_sizes_0"), val = tensor<int32, [2]>([32, 32])];
int32 var_475_axis_0 = const()[name = string("op_475_axis_0"), val = int32(-1)];
tensor<fp16, [1, 16, 1, 32]> var_475_0, tensor<fp16, [1, 16, 1, 32]> var_475_1 = split(axis = var_475_axis_0, split_sizes = var_475_split_sizes_0, x = q_1)[name = string("op_475")];
fp16 const_5_promoted = const()[name = string("const_5_promoted"), val = fp16(-0x1p+0)];
tensor<fp16, [1, 16, 1, 32]> var_477 = mul(x = var_475_1, y = const_5_promoted)[name = string("op_477")];
int32 var_479 = const()[name = string("op_479"), val = int32(-1)];
bool var_480_interleave_0 = const()[name = string("op_480_interleave_0"), val = bool(false)];
tensor<fp16, [1, 16, 1, 64]> var_480 = concat(axis = var_479, interleave = var_480_interleave_0, values = (var_477, var_475_0))[name = string("op_480")];
tensor<fp16, [1, 16, 1, 64]> var_481 = mul(x = var_480, y = sin)[name = string("op_481")];
tensor<fp16, [1, 16, 1, 64]> q = add(x = var_474, y = var_481)[name = string("q")];
tensor<fp16, [1, 8, 1, 64]> var_484 = mul(x = k_1, y = cos)[name = string("op_484")];
tensor<int32, [2]> var_485_split_sizes_0 = const()[name = string("op_485_split_sizes_0"), val = tensor<int32, [2]>([32, 32])];
int32 var_485_axis_0 = const()[name = string("op_485_axis_0"), val = int32(-1)];
tensor<fp16, [1, 8, 1, 32]> var_485_0, tensor<fp16, [1, 8, 1, 32]> var_485_1 = split(axis = var_485_axis_0, split_sizes = var_485_split_sizes_0, x = k_1)[name = string("op_485")];
fp16 const_6_promoted = const()[name = string("const_6_promoted"), val = fp16(-0x1p+0)];
tensor<fp16, [1, 8, 1, 32]> var_487 = mul(x = var_485_1, y = const_6_promoted)[name = string("op_487")];
int32 var_489 = const()[name = string("op_489"), val = int32(-1)];
bool var_490_interleave_0 = const()[name = string("op_490_interleave_0"), val = bool(false)];
tensor<fp16, [1, 8, 1, 64]> var_490 = concat(axis = var_489, interleave = var_490_interleave_0, values = (var_487, var_485_0))[name = string("op_490")];
tensor<fp16, [1, 8, 1, 64]> var_491 = mul(x = var_490, y = sin)[name = string("op_491")];
tensor<fp16, [1, 8, 1, 64]> k = add(x = var_484, y = var_491)[name = string("k")];
tensor<int32, [5]> K_cache_begin_0 = const()[name = string("K_cache_begin_0"), val = tensor<int32, [5]>([0, 0, 0, 0, 0])];
tensor<int32, [5]> K_cache_end_0 = const()[name = string("K_cache_end_0"), val = tensor<int32, [5]>([1, 1, 512, 1, 1024])];
tensor<bool, [5]> K_cache_end_mask_0 = const()[name = string("K_cache_end_mask_0"), val = tensor<bool, [5]>([false, true, true, true, true])];
tensor<bool, [5]> K_cache_squeeze_mask_0 = const()[name = string("K_cache_squeeze_mask_0"), val = tensor<bool, [5]>([true, false, false, false, false])];
tensor<fp16, [1, 512, 1, 1024]> K_cache_cast_fp16 = slice_by_index(begin = K_cache_begin_0, end = K_cache_end_0, end_mask = K_cache_end_mask_0, squeeze_mask = K_cache_squeeze_mask_0, x = kv_cache_in)[name = string("K_cache_cast_fp16")];
tensor<int32, [5]> V_cache_begin_0 = const()[name = string("V_cache_begin_0"), val = tensor<int32, [5]>([1, 0, 0, 0, 0])];
tensor<int32, [5]> V_cache_end_0 = const()[name = string("V_cache_end_0"), val = tensor<int32, [5]>([2, 1, 512, 1, 1024])];
tensor<bool, [5]> V_cache_end_mask_0 = const()[name = string("V_cache_end_mask_0"), val = tensor<bool, [5]>([false, true, true, true, true])];
tensor<bool, [5]> V_cache_squeeze_mask_0 = const()[name = string("V_cache_squeeze_mask_0"), val = tensor<bool, [5]>([true, false, false, false, false])];
tensor<fp16, [1, 512, 1, 1024]> V_cache_cast_fp16 = slice_by_index(begin = V_cache_begin_0, end = V_cache_end_0, end_mask = V_cache_end_mask_0, squeeze_mask = V_cache_squeeze_mask_0, x = kv_cache_in)[name = string("V_cache_cast_fp16")];
tensor<int32, [4]> var_504 = const()[name = string("op_504"), val = tensor<int32, [4]>([0, 1, 3, 2])];
tensor<int32, [4]> var_510 = const()[name = string("op_510"), val = tensor<int32, [4]>([1, 1024, 1, 1])];
tensor<fp16, [1, 16, 64, 1]> var_505 = transpose(perm = var_504, x = q)[name = string("transpose_13")];
tensor<fp16, [1, 1024, 1, 1]> query = reshape(shape = var_510, x = var_505)[name = string("query")];
tensor<int32, [4]> var_516 = const()[name = string("op_516"), val = tensor<int32, [4]>([0, 1, 3, 2])];
tensor<int32, [4]> var_522 = const()[name = string("op_522"), val = tensor<int32, [4]>([1, 512, 1, 1])];
tensor<fp16, [1, 8, 64, 1]> var_517 = transpose(perm = var_516, x = k)[name = string("transpose_12")];
tensor<fp16, [1, 512, 1, 1]> k_slice = reshape(shape = var_522, x = var_517)[name = string("k_slice")];
int32 var_537 = const()[name = string("op_537"), val = int32(-1)];
bool key_interleave_0 = const()[name = string("key_interleave_0"), val = bool(false)];
tensor<fp16, [1, 512, 1, 1025]> key_cast_fp16 = concat(axis = var_537, interleave = key_interleave_0, values = (K_cache_cast_fp16, k_slice))[name = string("key_cast_fp16")];
int32 var_540 = const()[name = string("op_540"), val = int32(-1)];
bool var_541_interleave_0 = const()[name = string("op_541_interleave_0"), val = bool(false)];
tensor<fp16, [1, 512, 1, 1025]> var_541_cast_fp16 = concat(axis = var_540, interleave = var_541_interleave_0, values = (V_cache_cast_fp16, var_428))[name = string("op_541_cast_fp16")];
tensor<fp16, [1, 1024, 1, 1]> var_542 = mul(x = query, y = attn_scale)[name = string("op_542")];
tensor<int32, [16]> tile_0 = const()[name = string("tile_0"), val = tensor<int32, [16]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(52346496)))];
int32 var_545_axis_0 = const()[name = string("op_545_axis_0"), val = int32(1)];
tensor<fp16, [1, 64, 1, 1]> var_545_0, tensor<fp16, [1, 64, 1, 1]> var_545_1, tensor<fp16, [1, 64, 1, 1]> var_545_2, tensor<fp16, [1, 64, 1, 1]> var_545_3, tensor<fp16, [1, 64, 1, 1]> var_545_4, tensor<fp16, [1, 64, 1, 1]> var_545_5, tensor<fp16, [1, 64, 1, 1]> var_545_6, tensor<fp16, [1, 64, 1, 1]> var_545_7, tensor<fp16, [1, 64, 1, 1]> var_545_8, tensor<fp16, [1, 64, 1, 1]> var_545_9, tensor<fp16, [1, 64, 1, 1]> var_545_10, tensor<fp16, [1, 64, 1, 1]> var_545_11, tensor<fp16, [1, 64, 1, 1]> var_545_12, tensor<fp16, [1, 64, 1, 1]> var_545_13, tensor<fp16, [1, 64, 1, 1]> var_545_14, tensor<fp16, [1, 64, 1, 1]> var_545_15 = split(axis = var_545_axis_0, split_sizes = tile_0, x = var_542)[name = string("op_545")];
tensor<int32, [4]> var_564_perm_0 = const()[name = string("op_564_perm_0"), val = tensor<int32, [4]>([0, 3, 2, 1])];
tensor<int32, [8]> tile_1 = const()[name = string("tile_1"), val = tensor<int32, [8]>([64, 64, 64, 64, 64, 64, 64, 64])];
int32 var_567_axis_0 = const()[name = string("op_567_axis_0"), val = int32(3)];
tensor<fp16, [1, 1025, 1, 512]> var_564_cast_fp16 = transpose(perm = var_564_perm_0, x = key_cast_fp16)[name = string("transpose_11")];
tensor<fp16, [1, 1025, 1, 64]> var_567_cast_fp16_0, tensor<fp16, [1, 1025, 1, 64]> var_567_cast_fp16_1, tensor<fp16, [1, 1025, 1, 64]> var_567_cast_fp16_2, tensor<fp16, [1, 1025, 1, 64]> var_567_cast_fp16_3, tensor<fp16, [1, 1025, 1, 64]> var_567_cast_fp16_4, tensor<fp16, [1, 1025, 1, 64]> var_567_cast_fp16_5, tensor<fp16, [1, 1025, 1, 64]> var_567_cast_fp16_6, tensor<fp16, [1, 1025, 1, 64]> var_567_cast_fp16_7 = split(axis = var_567_axis_0, split_sizes = tile_1, x = var_564_cast_fp16)[name = string("op_567_cast_fp16")];
tensor<int32, [8]> tile_2 = const()[name = string("tile_2"), val = tensor<int32, [8]>([64, 64, 64, 64, 64, 64, 64, 64])];
int32 var_578_axis_0 = const()[name = string("op_578_axis_0"), val = int32(1)];
tensor<fp16, [1, 64, 1, 1025]> var_578_cast_fp16_0, tensor<fp16, [1, 64, 1, 1025]> var_578_cast_fp16_1, tensor<fp16, [1, 64, 1, 1025]> var_578_cast_fp16_2, tensor<fp16, [1, 64, 1, 1025]> var_578_cast_fp16_3, tensor<fp16, [1, 64, 1, 1025]> var_578_cast_fp16_4, tensor<fp16, [1, 64, 1, 1025]> var_578_cast_fp16_5, tensor<fp16, [1, 64, 1, 1025]> var_578_cast_fp16_6, tensor<fp16, [1, 64, 1, 1025]> var_578_cast_fp16_7 = split(axis = var_578_axis_0, split_sizes = tile_2, x = var_541_cast_fp16)[name = string("op_578_cast_fp16")];
string scores_1_equation_0 = const()[name = string("scores_1_equation_0"), val = string("bkhc,bchq->bkhq")];
tensor<fp16, [1, 1025, 1, 1]> scores_1_cast_fp16 = einsum(equation = scores_1_equation_0, values = (var_567_cast_fp16_0, var_545_0))[name = string("scores_1_cast_fp16")];
tensor<fp16, [1, 1025, 1, 1]> var_592_cast_fp16 = add(x = scores_1_cast_fp16, y = causal_mask)[name = string("op_592_cast_fp16")];
int32 var_593 = const()[name = string("op_593"), val = int32(1)];
tensor<fp16, [1, 1025, 1, 1]> var_595_cast_fp16 = softmax(axis = var_593, x = var_592_cast_fp16)[name = string("op_595_cast_fp16")];
string var_599_equation_0 = const()[name = string("op_599_equation_0"), val = string("bchk,bkhq->bchq")];
tensor<fp16, [1, 64, 1, 1]> var_599_cast_fp16 = einsum(equation = var_599_equation_0, values = (var_578_cast_fp16_0, var_595_cast_fp16))[name = string("op_599_cast_fp16")];
string scores_3_equation_0 = const()[name = string("scores_3_equation_0"), val = string("bkhc,bchq->bkhq")];
tensor<fp16, [1, 1025, 1, 1]> scores_3_cast_fp16 = einsum(equation = scores_3_equation_0, values = (var_567_cast_fp16_0, var_545_1))[name = string("scores_3_cast_fp16")];
tensor<fp16, [1, 1025, 1, 1]> var_605_cast_fp16 = add(x = scores_3_cast_fp16, y = causal_mask)[name = string("op_605_cast_fp16")];
int32 var_606 = const()[name = string("op_606"), val = int32(1)];
tensor<fp16, [1, 1025, 1, 1]> var_608_cast_fp16 = softmax(axis = var_606, x = var_605_cast_fp16)[name = string("op_608_cast_fp16")];
string var_612_equation_0 = const()[name = string("op_612_equation_0"), val = string("bchk,bkhq->bchq")];
tensor<fp16, [1, 64, 1, 1]> var_612_cast_fp16 = einsum(equation = var_612_equation_0, values = (var_578_cast_fp16_0, var_608_cast_fp16))[name = string("op_612_cast_fp16")];
string scores_5_equation_0 = const()[name = string("scores_5_equation_0"), val = string("bkhc,bchq->bkhq")];
tensor<fp16, [1, 1025, 1, 1]> scores_5_cast_fp16 = einsum(equation = scores_5_equation_0, values = (var_567_cast_fp16_1, var_545_2))[name = string("scores_5_cast_fp16")];
tensor<fp16, [1, 1025, 1, 1]> var_618_cast_fp16 = add(x = scores_5_cast_fp16, y = causal_mask)[name = string("op_618_cast_fp16")];
int32 var_619 = const()[name = string("op_619"), val = int32(1)];
tensor<fp16, [1, 1025, 1, 1]> var_621_cast_fp16 = softmax(axis = var_619, x = var_618_cast_fp16)[name = string("op_621_cast_fp16")];
string var_625_equation_0 = const()[name = string("op_625_equation_0"), val = string("bchk,bkhq->bchq")];
tensor<fp16, [1, 64, 1, 1]> var_625_cast_fp16 = einsum(equation = var_625_equation_0, values = (var_578_cast_fp16_1, var_621_cast_fp16))[name = string("op_625_cast_fp16")];
string scores_7_equation_0 = const()[name = string("scores_7_equation_0"), val = string("bkhc,bchq->bkhq")];
tensor<fp16, [1, 1025, 1, 1]> scores_7_cast_fp16 = einsum(equation = scores_7_equation_0, values = (var_567_cast_fp16_1, var_545_3))[name = string("scores_7_cast_fp16")];
tensor<fp16, [1, 1025, 1, 1]> var_631_cast_fp16 = add(x = scores_7_cast_fp16, y = causal_mask)[name = string("op_631_cast_fp16")];
int32 var_632 = const()[name = string("op_632"), val = int32(1)];
tensor<fp16, [1, 1025, 1, 1]> var_634_cast_fp16 = softmax(axis = var_632, x = var_631_cast_fp16)[name = string("op_634_cast_fp16")];
string var_638_equation_0 = const()[name = string("op_638_equation_0"), val = string("bchk,bkhq->bchq")];
tensor<fp16, [1, 64, 1, 1]> var_638_cast_fp16 = einsum(equation = var_638_equation_0, values = (var_578_cast_fp16_1, var_634_cast_fp16))[name = string("op_638_cast_fp16")];
string scores_9_equation_0 = const()[name = string("scores_9_equation_0"), val = string("bkhc,bchq->bkhq")];
tensor<fp16, [1, 1025, 1, 1]> scores_9_cast_fp16 = einsum(equation = scores_9_equation_0, values = (var_567_cast_fp16_2, var_545_4))[name = string("scores_9_cast_fp16")];
tensor<fp16, [1, 1025, 1, 1]> var_644_cast_fp16 = add(x = scores_9_cast_fp16, y = causal_mask)[name = string("op_644_cast_fp16")];
int32 var_645 = const()[name = string("op_645"), val = int32(1)];
tensor<fp16, [1, 1025, 1, 1]> var_647_cast_fp16 = softmax(axis = var_645, x = var_644_cast_fp16)[name = string("op_647_cast_fp16")];
string var_651_equation_0 = const()[name = string("op_651_equation_0"), val = string("bchk,bkhq->bchq")];
tensor<fp16, [1, 64, 1, 1]> var_651_cast_fp16 = einsum(equation = var_651_equation_0, values = (var_578_cast_fp16_2, var_647_cast_fp16))[name = string("op_651_cast_fp16")];
string scores_11_equation_0 = const()[name = string("scores_11_equation_0"), val = string("bkhc,bchq->bkhq")];
tensor<fp16, [1, 1025, 1, 1]> scores_11_cast_fp16 = einsum(equation = scores_11_equation_0, values = (var_567_cast_fp16_2, var_545_5))[name = string("scores_11_cast_fp16")];
tensor<fp16, [1, 1025, 1, 1]> var_657_cast_fp16 = add(x = scores_11_cast_fp16, y = causal_mask)[name = string("op_657_cast_fp16")];
int32 var_658 = const()[name = string("op_658"), val = int32(1)];
tensor<fp16, [1, 1025, 1, 1]> var_660_cast_fp16 = softmax(axis = var_658, x = var_657_cast_fp16)[name = string("op_660_cast_fp16")];
string var_664_equation_0 = const()[name = string("op_664_equation_0"), val = string("bchk,bkhq->bchq")];
tensor<fp16, [1, 64, 1, 1]> var_664_cast_fp16 = einsum(equation = var_664_equation_0, values = (var_578_cast_fp16_2, var_660_cast_fp16))[name = string("op_664_cast_fp16")];
string scores_13_equation_0 = const()[name = string("scores_13_equation_0"), val = string("bkhc,bchq->bkhq")];
tensor<fp16, [1, 1025, 1, 1]> scores_13_cast_fp16 = einsum(equation = scores_13_equation_0, values = (var_567_cast_fp16_3, var_545_6))[name = string("scores_13_cast_fp16")];
tensor<fp16, [1, 1025, 1, 1]> var_670_cast_fp16 = add(x = scores_13_cast_fp16, y = causal_mask)[name = string("op_670_cast_fp16")];
int32 var_671 = const()[name = string("op_671"), val = int32(1)];
tensor<fp16, [1, 1025, 1, 1]> var_673_cast_fp16 = softmax(axis = var_671, x = var_670_cast_fp16)[name = string("op_673_cast_fp16")];
string var_677_equation_0 = const()[name = string("op_677_equation_0"), val = string("bchk,bkhq->bchq")];
tensor<fp16, [1, 64, 1, 1]> var_677_cast_fp16 = einsum(equation = var_677_equation_0, values = (var_578_cast_fp16_3, var_673_cast_fp16))[name = string("op_677_cast_fp16")];
string scores_15_equation_0 = const()[name = string("scores_15_equation_0"), val = string("bkhc,bchq->bkhq")];
tensor<fp16, [1, 1025, 1, 1]> scores_15_cast_fp16 = einsum(equation = scores_15_equation_0, values = (var_567_cast_fp16_3, var_545_7))[name = string("scores_15_cast_fp16")];
tensor<fp16, [1, 1025, 1, 1]> var_683_cast_fp16 = add(x = scores_15_cast_fp16, y = causal_mask)[name = string("op_683_cast_fp16")];
int32 var_684 = const()[name = string("op_684"), val = int32(1)];
tensor<fp16, [1, 1025, 1, 1]> var_686_cast_fp16 = softmax(axis = var_684, x = var_683_cast_fp16)[name = string("op_686_cast_fp16")];
string var_690_equation_0 = const()[name = string("op_690_equation_0"), val = string("bchk,bkhq->bchq")];
tensor<fp16, [1, 64, 1, 1]> var_690_cast_fp16 = einsum(equation = var_690_equation_0, values = (var_578_cast_fp16_3, var_686_cast_fp16))[name = string("op_690_cast_fp16")];
string scores_17_equation_0 = const()[name = string("scores_17_equation_0"), val = string("bkhc,bchq->bkhq")];
tensor<fp16, [1, 1025, 1, 1]> scores_17_cast_fp16 = einsum(equation = scores_17_equation_0, values = (var_567_cast_fp16_4, var_545_8))[name = string("scores_17_cast_fp16")];
tensor<fp16, [1, 1025, 1, 1]> var_696_cast_fp16 = add(x = scores_17_cast_fp16, y = causal_mask)[name = string("op_696_cast_fp16")];
int32 var_697 = const()[name = string("op_697"), val = int32(1)];
tensor<fp16, [1, 1025, 1, 1]> var_699_cast_fp16 = softmax(axis = var_697, x = var_696_cast_fp16)[name = string("op_699_cast_fp16")];
string var_703_equation_0 = const()[name = string("op_703_equation_0"), val = string("bchk,bkhq->bchq")];
tensor<fp16, [1, 64, 1, 1]> var_703_cast_fp16 = einsum(equation = var_703_equation_0, values = (var_578_cast_fp16_4, var_699_cast_fp16))[name = string("op_703_cast_fp16")];
string scores_19_equation_0 = const()[name = string("scores_19_equation_0"), val = string("bkhc,bchq->bkhq")];
tensor<fp16, [1, 1025, 1, 1]> scores_19_cast_fp16 = einsum(equation = scores_19_equation_0, values = (var_567_cast_fp16_4, var_545_9))[name = string("scores_19_cast_fp16")];
tensor<fp16, [1, 1025, 1, 1]> var_709_cast_fp16 = add(x = scores_19_cast_fp16, y = causal_mask)[name = string("op_709_cast_fp16")];
int32 var_710 = const()[name = string("op_710"), val = int32(1)];
tensor<fp16, [1, 1025, 1, 1]> var_712_cast_fp16 = softmax(axis = var_710, x = var_709_cast_fp16)[name = string("op_712_cast_fp16")];
string var_716_equation_0 = const()[name = string("op_716_equation_0"), val = string("bchk,bkhq->bchq")];
tensor<fp16, [1, 64, 1, 1]> var_716_cast_fp16 = einsum(equation = var_716_equation_0, values = (var_578_cast_fp16_4, var_712_cast_fp16))[name = string("op_716_cast_fp16")];
string scores_21_equation_0 = const()[name = string("scores_21_equation_0"), val = string("bkhc,bchq->bkhq")];
tensor<fp16, [1, 1025, 1, 1]> scores_21_cast_fp16 = einsum(equation = scores_21_equation_0, values = (var_567_cast_fp16_5, var_545_10))[name = string("scores_21_cast_fp16")];
tensor<fp16, [1, 1025, 1, 1]> var_722_cast_fp16 = add(x = scores_21_cast_fp16, y = causal_mask)[name = string("op_722_cast_fp16")];
int32 var_723 = const()[name = string("op_723"), val = int32(1)];
tensor<fp16, [1, 1025, 1, 1]> var_725_cast_fp16 = softmax(axis = var_723, x = var_722_cast_fp16)[name = string("op_725_cast_fp16")];
string var_729_equation_0 = const()[name = string("op_729_equation_0"), val = string("bchk,bkhq->bchq")];
tensor<fp16, [1, 64, 1, 1]> var_729_cast_fp16 = einsum(equation = var_729_equation_0, values = (var_578_cast_fp16_5, var_725_cast_fp16))[name = string("op_729_cast_fp16")];
string scores_23_equation_0 = const()[name = string("scores_23_equation_0"), val = string("bkhc,bchq->bkhq")];
tensor<fp16, [1, 1025, 1, 1]> scores_23_cast_fp16 = einsum(equation = scores_23_equation_0, values = (var_567_cast_fp16_5, var_545_11))[name = string("scores_23_cast_fp16")];
tensor<fp16, [1, 1025, 1, 1]> var_735_cast_fp16 = add(x = scores_23_cast_fp16, y = causal_mask)[name = string("op_735_cast_fp16")];
int32 var_736 = const()[name = string("op_736"), val = int32(1)];
tensor<fp16, [1, 1025, 1, 1]> var_738_cast_fp16 = softmax(axis = var_736, x = var_735_cast_fp16)[name = string("op_738_cast_fp16")];
string var_742_equation_0 = const()[name = string("op_742_equation_0"), val = string("bchk,bkhq->bchq")];
tensor<fp16, [1, 64, 1, 1]> var_742_cast_fp16 = einsum(equation = var_742_equation_0, values = (var_578_cast_fp16_5, var_738_cast_fp16))[name = string("op_742_cast_fp16")];
string scores_25_equation_0 = const()[name = string("scores_25_equation_0"), val = string("bkhc,bchq->bkhq")];
tensor<fp16, [1, 1025, 1, 1]> scores_25_cast_fp16 = einsum(equation = scores_25_equation_0, values = (var_567_cast_fp16_6, var_545_12))[name = string("scores_25_cast_fp16")];
tensor<fp16, [1, 1025, 1, 1]> var_748_cast_fp16 = add(x = scores_25_cast_fp16, y = causal_mask)[name = string("op_748_cast_fp16")];
int32 var_749 = const()[name = string("op_749"), val = int32(1)];
tensor<fp16, [1, 1025, 1, 1]> var_751_cast_fp16 = softmax(axis = var_749, x = var_748_cast_fp16)[name = string("op_751_cast_fp16")];
string var_755_equation_0 = const()[name = string("op_755_equation_0"), val = string("bchk,bkhq->bchq")];
tensor<fp16, [1, 64, 1, 1]> var_755_cast_fp16 = einsum(equation = var_755_equation_0, values = (var_578_cast_fp16_6, var_751_cast_fp16))[name = string("op_755_cast_fp16")];
string scores_27_equation_0 = const()[name = string("scores_27_equation_0"), val = string("bkhc,bchq->bkhq")];
tensor<fp16, [1, 1025, 1, 1]> scores_27_cast_fp16 = einsum(equation = scores_27_equation_0, values = (var_567_cast_fp16_6, var_545_13))[name = string("scores_27_cast_fp16")];
tensor<fp16, [1, 1025, 1, 1]> var_761_cast_fp16 = add(x = scores_27_cast_fp16, y = causal_mask)[name = string("op_761_cast_fp16")];
int32 var_762 = const()[name = string("op_762"), val = int32(1)];
tensor<fp16, [1, 1025, 1, 1]> var_764_cast_fp16 = softmax(axis = var_762, x = var_761_cast_fp16)[name = string("op_764_cast_fp16")];
string var_768_equation_0 = const()[name = string("op_768_equation_0"), val = string("bchk,bkhq->bchq")];
tensor<fp16, [1, 64, 1, 1]> var_768_cast_fp16 = einsum(equation = var_768_equation_0, values = (var_578_cast_fp16_6, var_764_cast_fp16))[name = string("op_768_cast_fp16")];
string scores_29_equation_0 = const()[name = string("scores_29_equation_0"), val = string("bkhc,bchq->bkhq")];
tensor<fp16, [1, 1025, 1, 1]> scores_29_cast_fp16 = einsum(equation = scores_29_equation_0, values = (var_567_cast_fp16_7, var_545_14))[name = string("scores_29_cast_fp16")];
tensor<fp16, [1, 1025, 1, 1]> var_774_cast_fp16 = add(x = scores_29_cast_fp16, y = causal_mask)[name = string("op_774_cast_fp16")];
int32 var_775 = const()[name = string("op_775"), val = int32(1)];
tensor<fp16, [1, 1025, 1, 1]> var_777_cast_fp16 = softmax(axis = var_775, x = var_774_cast_fp16)[name = string("op_777_cast_fp16")];
string var_781_equation_0 = const()[name = string("op_781_equation_0"), val = string("bchk,bkhq->bchq")];
tensor<fp16, [1, 64, 1, 1]> var_781_cast_fp16 = einsum(equation = var_781_equation_0, values = (var_578_cast_fp16_7, var_777_cast_fp16))[name = string("op_781_cast_fp16")];
string scores_equation_0 = const()[name = string("scores_equation_0"), val = string("bkhc,bchq->bkhq")];
tensor<fp16, [1, 1025, 1, 1]> scores_cast_fp16 = einsum(equation = scores_equation_0, values = (var_567_cast_fp16_7, var_545_15))[name = string("scores_cast_fp16")];
tensor<fp16, [1, 1025, 1, 1]> var_787_cast_fp16 = add(x = scores_cast_fp16, y = causal_mask)[name = string("op_787_cast_fp16")];
int32 var_788 = const()[name = string("op_788"), val = int32(1)];
tensor<fp16, [1, 1025, 1, 1]> var_790_cast_fp16 = softmax(axis = var_788, x = var_787_cast_fp16)[name = string("op_790_cast_fp16")];
string var_794_equation_0 = const()[name = string("op_794_equation_0"), val = string("bchk,bkhq->bchq")];
tensor<fp16, [1, 64, 1, 1]> var_794_cast_fp16 = einsum(equation = var_794_equation_0, values = (var_578_cast_fp16_7, var_790_cast_fp16))[name = string("op_794_cast_fp16")];
int32 var_796 = const()[name = string("op_796"), val = int32(1)];
bool input_23_interleave_0 = const()[name = string("input_23_interleave_0"), val = bool(false)];
tensor<fp16, [1, 1024, 1, 1]> input_23_cast_fp16 = concat(axis = var_796, interleave = input_23_interleave_0, values = (var_599_cast_fp16, var_612_cast_fp16, var_625_cast_fp16, var_638_cast_fp16, var_651_cast_fp16, var_664_cast_fp16, var_677_cast_fp16, var_690_cast_fp16, var_703_cast_fp16, var_716_cast_fp16, var_729_cast_fp16, var_742_cast_fp16, var_755_cast_fp16, var_768_cast_fp16, var_781_cast_fp16, var_794_cast_fp16))[name = string("input_23_cast_fp16")];
string out_pad_type_0 = const()[name = string("out_pad_type_0"), val = string("valid")];
tensor<int32, [2]> out_strides_0 = const()[name = string("out_strides_0"), val = tensor<int32, [2]>([1, 1])];
tensor<int32, [4]> out_pad_0 = const()[name = string("out_pad_0"), val = tensor<int32, [4]>([0, 0, 0, 0])];
tensor<int32, [2]> out_dilations_0 = const()[name = string("out_dilations_0"), val = tensor<int32, [2]>([1, 1])];
int32 out_groups_0 = const()[name = string("out_groups_0"), val = int32(1)];
tensor<fp16, [1024, 1024, 1, 1]> layers_1_self_attn_out_proj_weight_promoted_to_fp16_palettized = constexpr_lut_to_dense(indices = tensor<uint6, [1024, 1024, 1, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(52346624))), lut = tensor<fp16, [32, 1, 1, 1, 64, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(53133120))))[name = string("layers_1_self_attn_out_proj_weight_promoted_to_fp16_palettized")];
tensor<fp16, [1, 1024, 1, 1]> out_cast_fp16 = conv(dilations = out_dilations_0, groups = out_groups_0, pad = out_pad_0, pad_type = out_pad_type_0, strides = out_strides_0, weight = layers_1_self_attn_out_proj_weight_promoted_to_fp16_palettized, x = input_23_cast_fp16)[name = string("out_cast_fp16")];
tensor<int32, [1]> var_810_axes_0 = const()[name = string("op_810_axes_0"), val = tensor<int32, [1]>([2])];
tensor<fp16, [1, 1024, 1]> var_810_cast_fp16 = squeeze(axes = var_810_axes_0, x = out_cast_fp16)[name = string("op_810_cast_fp16")];
tensor<int32, [3]> var_814 = const()[name = string("op_814"), val = tensor<int32, [3]>([0, 2, 1])];
tensor<fp16, [1, 1, 1024]> op_out_3_cast_fp16 = transpose(perm = var_814, x = var_810_cast_fp16)[name = string("transpose_10")];
tensor<fp16, [1, 1, 1024]> x_11_cast_fp16 = add(x = x_5_cast_fp16, y = op_out_3_cast_fp16)[name = string("x_11_cast_fp16")];
fp16 const_13_promoted_to_fp16 = const()[name = string("const_13_promoted_to_fp16"), val = fp16(-0x1p+0)];
tensor<fp16, [1, 1, 1024]> var_818_cast_fp16 = mul(x = x_11_cast_fp16, y = const_13_promoted_to_fp16)[name = string("op_818_cast_fp16")];
int32 var_820 = const()[name = string("op_820"), val = int32(-1)];
bool input_25_interleave_0 = const()[name = string("input_25_interleave_0"), val = bool(false)];
tensor<fp16, [1, 1, 2048]> input_25_cast_fp16 = concat(axis = var_820, interleave = input_25_interleave_0, values = (x_11_cast_fp16, var_818_cast_fp16))[name = string("input_25_cast_fp16")];
tensor<int32, [1]> normed_13_axes_0 = const()[name = string("normed_13_axes_0"), val = tensor<int32, [1]>([-1])];
fp16 var_826_to_fp16 = const()[name = string("op_826_to_fp16"), val = fp16(0x1.5p-17)];
tensor<fp16, [1, 1, 2048]> normed_13_cast_fp16 = layer_norm(axes = normed_13_axes_0, epsilon = var_826_to_fp16, x = input_25_cast_fp16)[name = string("normed_13_cast_fp16")];
tensor<int32, [2]> var_829_split_sizes_0 = const()[name = string("op_829_split_sizes_0"), val = tensor<int32, [2]>([1024, 1024])];
int32 var_829_axis_0 = const()[name = string("op_829_axis_0"), val = int32(-1)];
tensor<fp16, [1, 1, 1024]> var_829_cast_fp16_0, tensor<fp16, [1, 1, 1024]> var_829_cast_fp16_1 = split(axis = var_829_axis_0, split_sizes = var_829_split_sizes_0, x = normed_13_cast_fp16)[name = string("op_829_cast_fp16")];
tensor<fp16, [1024]> layers_1_ffn_norm_weight_promoted_to_fp16 = const()[name = string("layers_1_ffn_norm_weight_promoted_to_fp16"), val = tensor<fp16, [1024]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(53137280)))];
tensor<fp16, [1, 1, 1024]> normed_15_cast_fp16 = mul(x = var_829_cast_fp16_0, y = layers_1_ffn_norm_weight_promoted_to_fp16)[name = string("normed_15_cast_fp16")];
tensor<int32, [3]> var_835 = const()[name = string("op_835"), val = tensor<int32, [3]>([0, 2, 1])];
tensor<int32, [1]> var_838_axes_0 = const()[name = string("op_838_axes_0"), val = tensor<int32, [1]>([2])];
tensor<fp16, [1, 1024, 1]> var_836_cast_fp16 = transpose(perm = var_835, x = normed_15_cast_fp16)[name = string("transpose_9")];
tensor<fp16, [1, 1024, 1, 1]> var_838_cast_fp16 = expand_dims(axes = var_838_axes_0, x = var_836_cast_fp16)[name = string("op_838_cast_fp16")];
string input_29_pad_type_0 = const()[name = string("input_29_pad_type_0"), val = string("valid")];
tensor<int32, [2]> input_29_strides_0 = const()[name = string("input_29_strides_0"), val = tensor<int32, [2]>([1, 1])];
tensor<int32, [4]> input_29_pad_0 = const()[name = string("input_29_pad_0"), val = tensor<int32, [4]>([0, 0, 0, 0])];
tensor<int32, [2]> input_29_dilations_0 = const()[name = string("input_29_dilations_0"), val = tensor<int32, [2]>([1, 1])];
int32 input_29_groups_0 = const()[name = string("input_29_groups_0"), val = int32(1)];
tensor<fp16, [1, 4608, 1, 1]> input_29 = conv(dilations = input_29_dilations_0, groups = input_29_groups_0, pad = input_29_pad_0, pad_type = input_29_pad_type_0, strides = input_29_strides_0, weight = layers_1_feed_forward_w1_weight_palettized, x = var_838_cast_fp16)[name = string("input_29")];
string b_3_pad_type_0 = const()[name = string("b_3_pad_type_0"), val = string("valid")];
tensor<int32, [2]> b_3_strides_0 = const()[name = string("b_3_strides_0"), val = tensor<int32, [2]>([1, 1])];
tensor<int32, [4]> b_3_pad_0 = const()[name = string("b_3_pad_0"), val = tensor<int32, [4]>([0, 0, 0, 0])];
tensor<int32, [2]> b_3_dilations_0 = const()[name = string("b_3_dilations_0"), val = tensor<int32, [2]>([1, 1])];
int32 b_3_groups_0 = const()[name = string("b_3_groups_0"), val = int32(1)];
tensor<fp16, [1, 4608, 1, 1]> b_3 = conv(dilations = b_3_dilations_0, groups = b_3_groups_0, pad = b_3_pad_0, pad_type = b_3_pad_type_0, strides = b_3_strides_0, weight = layers_1_feed_forward_w3_weight_palettized, x = var_838_cast_fp16)[name = string("b_3")];
tensor<fp16, [1, 4608, 1, 1]> var_866 = silu(x = input_29)[name = string("op_866")];
tensor<fp16, [1, 4608, 1, 1]> input_31 = mul(x = var_866, y = b_3)[name = string("input_31")];
string mlp_5_pad_type_0 = const()[name = string("mlp_5_pad_type_0"), val = string("valid")];
tensor<int32, [2]> mlp_5_strides_0 = const()[name = string("mlp_5_strides_0"), val = tensor<int32, [2]>([1, 1])];
tensor<int32, [4]> mlp_5_pad_0 = const()[name = string("mlp_5_pad_0"), val = tensor<int32, [4]>([0, 0, 0, 0])];
tensor<int32, [2]> mlp_5_dilations_0 = const()[name = string("mlp_5_dilations_0"), val = tensor<int32, [2]>([1, 1])];
int32 mlp_5_groups_0 = const()[name = string("mlp_5_groups_0"), val = int32(1)];
tensor<fp16, [1, 1024, 1, 1]> mlp_5 = conv(dilations = mlp_5_dilations_0, groups = mlp_5_groups_0, pad = mlp_5_pad_0, pad_type = mlp_5_pad_type_0, strides = mlp_5_strides_0, weight = layers_1_feed_forward_w2_weight_palettized, x = input_31)[name = string("mlp_5")];
tensor<int32, [1]> var_880_axes_0 = const()[name = string("op_880_axes_0"), val = tensor<int32, [1]>([2])];
tensor<fp16, [1, 1024, 1]> var_880 = squeeze(axes = var_880_axes_0, x = mlp_5)[name = string("op_880")];
tensor<int32, [3]> var_884 = const()[name = string("op_884"), val = tensor<int32, [3]>([0, 2, 1])];
tensor<fp16, [1, 1, 1024]> mlp_7 = transpose(perm = var_884, x = var_880)[name = string("transpose_8")];
tensor<fp16, [1, 1, 1024]> x_13_cast_fp16 = add(x = x_11_cast_fp16, y = mlp_7)[name = string("x_13_cast_fp16")];
fp16 const_14_promoted_to_fp16 = const()[name = string("const_14_promoted_to_fp16"), val = fp16(-0x1p+0)];
tensor<fp16, [1, 1, 1024]> var_888_cast_fp16 = mul(x = x_13_cast_fp16, y = const_14_promoted_to_fp16)[name = string("op_888_cast_fp16")];
int32 var_890 = const()[name = string("op_890"), val = int32(-1)];
bool input_33_interleave_0 = const()[name = string("input_33_interleave_0"), val = bool(false)];
tensor<fp16, [1, 1, 2048]> input_33_cast_fp16 = concat(axis = var_890, interleave = input_33_interleave_0, values = (x_13_cast_fp16, var_888_cast_fp16))[name = string("input_33_cast_fp16")];
tensor<int32, [1]> normed_17_axes_0 = const()[name = string("normed_17_axes_0"), val = tensor<int32, [1]>([-1])];
fp16 var_896_to_fp16 = const()[name = string("op_896_to_fp16"), val = fp16(0x1.5p-17)];
tensor<fp16, [1, 1, 2048]> normed_17_cast_fp16 = layer_norm(axes = normed_17_axes_0, epsilon = var_896_to_fp16, x = input_33_cast_fp16)[name = string("normed_17_cast_fp16")];
tensor<int32, [2]> var_899_split_sizes_0 = const()[name = string("op_899_split_sizes_0"), val = tensor<int32, [2]>([1024, 1024])];
int32 var_899_axis_0 = const()[name = string("op_899_axis_0"), val = int32(-1)];
tensor<fp16, [1, 1, 1024]> var_899_cast_fp16_0, tensor<fp16, [1, 1, 1024]> var_899_cast_fp16_1 = split(axis = var_899_axis_0, split_sizes = var_899_split_sizes_0, x = normed_17_cast_fp16)[name = string("op_899_cast_fp16")];
tensor<fp16, [1024]> layers_2_operator_norm_weight_promoted_to_fp16 = const()[name = string("layers_2_operator_norm_weight_promoted_to_fp16"), val = tensor<fp16, [1024]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(53139392)))];
tensor<fp16, [1, 1, 1024]> hidden_states_5_cast_fp16 = mul(x = var_899_cast_fp16_0, y = layers_2_operator_norm_weight_promoted_to_fp16)[name = string("hidden_states_5_cast_fp16")];
tensor<int32, [3]> var_905 = const()[name = string("op_905"), val = tensor<int32, [3]>([0, 2, 1])];
tensor<int32, [1]> var_908_axes_0 = const()[name = string("op_908_axes_0"), val = tensor<int32, [1]>([2])];
tensor<fp16, [1, 1024, 1]> var_906_cast_fp16 = transpose(perm = var_905, x = hidden_states_5_cast_fp16)[name = string("transpose_7")];
tensor<fp16, [1, 1024, 1, 1]> var_908_cast_fp16 = expand_dims(axes = var_908_axes_0, x = var_906_cast_fp16)[name = string("op_908_cast_fp16")];
string BCx_3_pad_type_0 = const()[name = string("BCx_3_pad_type_0"), val = string("valid")];
tensor<int32, [2]> BCx_3_strides_0 = const()[name = string("BCx_3_strides_0"), val = tensor<int32, [2]>([1, 1])];
tensor<int32, [4]> BCx_3_pad_0 = const()[name = string("BCx_3_pad_0"), val = tensor<int32, [4]>([0, 0, 0, 0])];
tensor<int32, [2]> BCx_3_dilations_0 = const()[name = string("BCx_3_dilations_0"), val = tensor<int32, [2]>([1, 1])];
int32 BCx_3_groups_0 = const()[name = string("BCx_3_groups_0"), val = int32(1)];
tensor<fp16, [1, 3072, 1, 1]> BCx_3 = conv(dilations = BCx_3_dilations_0, groups = BCx_3_groups_0, pad = BCx_3_pad_0, pad_type = BCx_3_pad_type_0, strides = BCx_3_strides_0, weight = layers_2_conv_in_proj_weight_palettized, x = var_908_cast_fp16)[name = string("BCx_3")];
tensor<int32, [3]> var_925_split_sizes_0 = const()[name = string("op_925_split_sizes_0"), val = tensor<int32, [3]>([1024, 1024, 1024])];
int32 var_925_axis_0 = const()[name = string("op_925_axis_0"), val = int32(1)];
tensor<fp16, [1, 1024, 1, 1]> var_925_0, tensor<fp16, [1, 1024, 1, 1]> var_925_1, tensor<fp16, [1, 1024, 1, 1]> var_925_2 = split(axis = var_925_axis_0, split_sizes = var_925_split_sizes_0, x = BCx_3)[name = string("op_925")];
tensor<fp16, [1, 1024, 1, 1]> Bx_3 = mul(x = var_925_0, y = var_925_2)[name = string("Bx_3")];
tensor<int32, [3]> var_931_begin_0 = const()[name = string("op_931_begin_0"), val = tensor<int32, [3]>([1, 0, 0])];
tensor<int32, [3]> var_931_end_0 = const()[name = string("op_931_end_0"), val = tensor<int32, [3]>([2, 1024, 3])];
tensor<bool, [3]> var_931_end_mask_0 = const()[name = string("op_931_end_mask_0"), val = tensor<bool, [3]>([false, true, true])];
tensor<bool, [3]> var_931_squeeze_mask_0 = const()[name = string("op_931_squeeze_mask_0"), val = tensor<bool, [3]>([true, false, false])];
tensor<fp16, [1024, 3]> var_931_cast_fp16 = slice_by_index(begin = var_931_begin_0, end = var_931_end_0, end_mask = var_931_end_mask_0, squeeze_mask = var_931_squeeze_mask_0, x = conv_state_in)[name = string("op_931_cast_fp16")];
tensor<int32, [1]> var_933_axes_0 = const()[name = string("op_933_axes_0"), val = tensor<int32, [1]>([0])];
tensor<fp16, [1, 1024, 3]> var_933_cast_fp16 = expand_dims(axes = var_933_axes_0, x = var_931_cast_fp16)[name = string("op_933_cast_fp16")];
tensor<int32, [1]> slot_3_axes_0 = const()[name = string("slot_3_axes_0"), val = tensor<int32, [1]>([2])];
tensor<fp16, [1, 1024, 1, 3]> slot_3_cast_fp16 = expand_dims(axes = slot_3_axes_0, x = var_933_cast_fp16)[name = string("slot_3_cast_fp16")];
tensor<int32, [4]> live_tail_3_begin_0 = const()[name = string("live_tail_3_begin_0"), val = tensor<int32, [4]>([0, 0, 0, 1])];
tensor<int32, [4]> live_tail_3_end_0 = const()[name = string("live_tail_3_end_0"), val = tensor<int32, [4]>([1, 1024, 1, 1])];
tensor<bool, [4]> live_tail_3_end_mask_0 = const()[name = string("live_tail_3_end_mask_0"), val = tensor<bool, [4]>([true, true, true, true])];
tensor<fp16, [1, 1024, 1, 2]> live_tail_3_cast_fp16 = slice_by_index(begin = live_tail_3_begin_0, end = live_tail_3_end_0, end_mask = live_tail_3_end_mask_0, x = slot_3_cast_fp16)[name = string("live_tail_3_cast_fp16")];
int32 var_942 = const()[name = string("op_942"), val = int32(-1)];
bool new_state_3_interleave_0 = const()[name = string("new_state_3_interleave_0"), val = bool(false)];
tensor<fp16, [1, 1024, 1, 3]> new_state_3_cast_fp16 = concat(axis = var_942, interleave = new_state_3_interleave_0, values = (live_tail_3_cast_fp16, Bx_3))[name = string("new_state_3_cast_fp16")];
tensor<int32, [1]> var_945_axes_0 = const()[name = string("op_945_axes_0"), val = tensor<int32, [1]>([0])];
tensor<fp16, [1024, 1, 3]> var_945_cast_fp16 = squeeze(axes = var_945_axes_0, x = new_state_3_cast_fp16)[name = string("op_945_cast_fp16")];
tensor<int32, [1]> var_947_axes_0 = const()[name = string("op_947_axes_0"), val = tensor<int32, [1]>([1])];
tensor<fp16, [1024, 3]> var_947_cast_fp16 = squeeze(axes = var_947_axes_0, x = var_945_cast_fp16)[name = string("op_947_cast_fp16")];
string conv_out_3_pad_type_0 = const()[name = string("conv_out_3_pad_type_0"), val = string("valid")];
int32 conv_out_3_groups_0 = const()[name = string("conv_out_3_groups_0"), val = int32(1024)];
tensor<int32, [2]> conv_out_3_strides_0 = const()[name = string("conv_out_3_strides_0"), val = tensor<int32, [2]>([1, 1])];
tensor<int32, [4]> conv_out_3_pad_0 = const()[name = string("conv_out_3_pad_0"), val = tensor<int32, [4]>([0, 0, 0, 0])];
tensor<int32, [2]> conv_out_3_dilations_0 = const()[name = string("conv_out_3_dilations_0"), val = tensor<int32, [2]>([1, 1])];
tensor<fp16, [1024, 1, 1, 3]> layers_2_conv_conv_weight_promoted_to_fp16_palettized = constexpr_lut_to_dense(indices = tensor<uint6, [1024, 1, 1, 3]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(53141504))), lut = tensor<fp16, [32, 1, 1, 1, 64, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(53143872))))[name = string("layers_2_conv_conv_weight_promoted_to_fp16_palettized")];
tensor<fp16, [1, 1024, 1, 1]> conv_out_3_cast_fp16 = conv(dilations = conv_out_3_dilations_0, groups = conv_out_3_groups_0, pad = conv_out_3_pad_0, pad_type = conv_out_3_pad_type_0, strides = conv_out_3_strides_0, weight = layers_2_conv_conv_weight_promoted_to_fp16_palettized, x = new_state_3_cast_fp16)[name = string("conv_out_3_cast_fp16")];
tensor<fp16, [1, 1024, 1, 1]> input_37_cast_fp16 = mul(x = var_925_1, y = conv_out_3_cast_fp16)[name = string("input_37_cast_fp16")];
string y_3_pad_type_0 = const()[name = string("y_3_pad_type_0"), val = string("valid")];
tensor<int32, [2]> y_3_strides_0 = const()[name = string("y_3_strides_0"), val = tensor<int32, [2]>([1, 1])];
tensor<int32, [4]> y_3_pad_0 = const()[name = string("y_3_pad_0"), val = tensor<int32, [4]>([0, 0, 0, 0])];
tensor<int32, [2]> y_3_dilations_0 = const()[name = string("y_3_dilations_0"), val = tensor<int32, [2]>([1, 1])];
int32 y_3_groups_0 = const()[name = string("y_3_groups_0"), val = int32(1)];
tensor<fp16, [1024, 1024, 1, 1]> layers_2_conv_out_proj_weight_promoted_to_fp16_palettized = constexpr_lut_to_dense(indices = tensor<uint6, [1024, 1024, 1, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(53148032))), lut = tensor<fp16, [32, 1, 1, 1, 64, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(53934528))))[name = string("layers_2_conv_out_proj_weight_promoted_to_fp16_palettized")];
tensor<fp16, [1, 1024, 1, 1]> y_3_cast_fp16 = conv(dilations = y_3_dilations_0, groups = y_3_groups_0, pad = y_3_pad_0, pad_type = y_3_pad_type_0, strides = y_3_strides_0, weight = layers_2_conv_out_proj_weight_promoted_to_fp16_palettized, x = input_37_cast_fp16)[name = string("y_3_cast_fp16")];
tensor<int32, [1]> var_973_axes_0 = const()[name = string("op_973_axes_0"), val = tensor<int32, [1]>([2])];
tensor<fp16, [1, 1024, 1]> var_973_cast_fp16 = squeeze(axes = var_973_axes_0, x = y_3_cast_fp16)[name = string("op_973_cast_fp16")];
tensor<int32, [3]> var_977 = const()[name = string("op_977"), val = tensor<int32, [3]>([0, 2, 1])];
tensor<fp16, [1, 1, 1024]> op_out_5_cast_fp16 = transpose(perm = var_977, x = var_973_cast_fp16)[name = string("transpose_6")];
tensor<fp16, [1, 1, 1024]> x_15_cast_fp16 = add(x = x_13_cast_fp16, y = op_out_5_cast_fp16)[name = string("x_15_cast_fp16")];
fp16 const_15_promoted_to_fp16 = const()[name = string("const_15_promoted_to_fp16"), val = fp16(-0x1p+0)];
tensor<fp16, [1, 1, 1024]> var_981_cast_fp16 = mul(x = x_15_cast_fp16, y = const_15_promoted_to_fp16)[name = string("op_981_cast_fp16")];
int32 var_983 = const()[name = string("op_983"), val = int32(-1)];
bool input_39_interleave_0 = const()[name = string("input_39_interleave_0"), val = bool(false)];
tensor<fp16, [1, 1, 2048]> input_39_cast_fp16 = concat(axis = var_983, interleave = input_39_interleave_0, values = (x_15_cast_fp16, var_981_cast_fp16))[name = string("input_39_cast_fp16")];
tensor<int32, [1]> normed_19_axes_0 = const()[name = string("normed_19_axes_0"), val = tensor<int32, [1]>([-1])];
fp16 var_989_to_fp16 = const()[name = string("op_989_to_fp16"), val = fp16(0x1.5p-17)];
tensor<fp16, [1, 1, 2048]> normed_19_cast_fp16 = layer_norm(axes = normed_19_axes_0, epsilon = var_989_to_fp16, x = input_39_cast_fp16)[name = string("normed_19_cast_fp16")];
tensor<int32, [2]> var_992_split_sizes_0 = const()[name = string("op_992_split_sizes_0"), val = tensor<int32, [2]>([1024, 1024])];
int32 var_992_axis_0 = const()[name = string("op_992_axis_0"), val = int32(-1)];
tensor<fp16, [1, 1, 1024]> var_992_cast_fp16_0, tensor<fp16, [1, 1, 1024]> var_992_cast_fp16_1 = split(axis = var_992_axis_0, split_sizes = var_992_split_sizes_0, x = normed_19_cast_fp16)[name = string("op_992_cast_fp16")];
tensor<fp16, [1024]> layers_2_ffn_norm_weight_promoted_to_fp16 = const()[name = string("layers_2_ffn_norm_weight_promoted_to_fp16"), val = tensor<fp16, [1024]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(53938688)))];
tensor<fp16, [1, 1, 1024]> normed_21_cast_fp16 = mul(x = var_992_cast_fp16_0, y = layers_2_ffn_norm_weight_promoted_to_fp16)[name = string("normed_21_cast_fp16")];
tensor<int32, [3]> var_998 = const()[name = string("op_998"), val = tensor<int32, [3]>([0, 2, 1])];
tensor<int32, [1]> var_1001_axes_0 = const()[name = string("op_1001_axes_0"), val = tensor<int32, [1]>([2])];
tensor<fp16, [1, 1024, 1]> var_999_cast_fp16 = transpose(perm = var_998, x = normed_21_cast_fp16)[name = string("transpose_5")];
tensor<fp16, [1, 1024, 1, 1]> var_1001_cast_fp16 = expand_dims(axes = var_1001_axes_0, x = var_999_cast_fp16)[name = string("op_1001_cast_fp16")];
string input_43_pad_type_0 = const()[name = string("input_43_pad_type_0"), val = string("valid")];
tensor<int32, [2]> input_43_strides_0 = const()[name = string("input_43_strides_0"), val = tensor<int32, [2]>([1, 1])];
tensor<int32, [4]> input_43_pad_0 = const()[name = string("input_43_pad_0"), val = tensor<int32, [4]>([0, 0, 0, 0])];
tensor<int32, [2]> input_43_dilations_0 = const()[name = string("input_43_dilations_0"), val = tensor<int32, [2]>([1, 1])];
int32 input_43_groups_0 = const()[name = string("input_43_groups_0"), val = int32(1)];
tensor<fp16, [1, 4608, 1, 1]> input_43 = conv(dilations = input_43_dilations_0, groups = input_43_groups_0, pad = input_43_pad_0, pad_type = input_43_pad_type_0, strides = input_43_strides_0, weight = layers_2_feed_forward_w1_weight_palettized, x = var_1001_cast_fp16)[name = string("input_43")];
string b_5_pad_type_0 = const()[name = string("b_5_pad_type_0"), val = string("valid")];
tensor<int32, [2]> b_5_strides_0 = const()[name = string("b_5_strides_0"), val = tensor<int32, [2]>([1, 1])];
tensor<int32, [4]> b_5_pad_0 = const()[name = string("b_5_pad_0"), val = tensor<int32, [4]>([0, 0, 0, 0])];
tensor<int32, [2]> b_5_dilations_0 = const()[name = string("b_5_dilations_0"), val = tensor<int32, [2]>([1, 1])];
int32 b_5_groups_0 = const()[name = string("b_5_groups_0"), val = int32(1)];
tensor<fp16, [1, 4608, 1, 1]> b_5 = conv(dilations = b_5_dilations_0, groups = b_5_groups_0, pad = b_5_pad_0, pad_type = b_5_pad_type_0, strides = b_5_strides_0, weight = layers_2_feed_forward_w3_weight_palettized, x = var_1001_cast_fp16)[name = string("b_5")];
tensor<fp16, [1, 4608, 1, 1]> var_1029 = silu(x = input_43)[name = string("op_1029")];
tensor<fp16, [1, 4608, 1, 1]> input_45 = mul(x = var_1029, y = b_5)[name = string("input_45")];
string mlp_9_pad_type_0 = const()[name = string("mlp_9_pad_type_0"), val = string("valid")];
tensor<int32, [2]> mlp_9_strides_0 = const()[name = string("mlp_9_strides_0"), val = tensor<int32, [2]>([1, 1])];
tensor<int32, [4]> mlp_9_pad_0 = const()[name = string("mlp_9_pad_0"), val = tensor<int32, [4]>([0, 0, 0, 0])];
tensor<int32, [2]> mlp_9_dilations_0 = const()[name = string("mlp_9_dilations_0"), val = tensor<int32, [2]>([1, 1])];
int32 mlp_9_groups_0 = const()[name = string("mlp_9_groups_0"), val = int32(1)];
tensor<fp16, [1, 1024, 1, 1]> mlp_9 = conv(dilations = mlp_9_dilations_0, groups = mlp_9_groups_0, pad = mlp_9_pad_0, pad_type = mlp_9_pad_type_0, strides = mlp_9_strides_0, weight = layers_2_feed_forward_w2_weight_palettized, x = input_45)[name = string("mlp_9")];
tensor<int32, [1]> var_1043_axes_0 = const()[name = string("op_1043_axes_0"), val = tensor<int32, [1]>([2])];
tensor<fp16, [1, 1024, 1]> var_1043 = squeeze(axes = var_1043_axes_0, x = mlp_9)[name = string("op_1043")];
tensor<int32, [3]> var_1047 = const()[name = string("op_1047"), val = tensor<int32, [3]>([0, 2, 1])];
tensor<fp16, [1, 1, 1024]> mlp_11 = transpose(perm = var_1047, x = var_1043)[name = string("transpose_4")];
tensor<fp16, [1, 1, 1024]> x_17_cast_fp16 = add(x = x_15_cast_fp16, y = mlp_11)[name = string("x_17_cast_fp16")];
fp16 const_16_promoted_to_fp16 = const()[name = string("const_16_promoted_to_fp16"), val = fp16(-0x1p+0)];
tensor<fp16, [1, 1, 1024]> var_1051_cast_fp16 = mul(x = x_17_cast_fp16, y = const_16_promoted_to_fp16)[name = string("op_1051_cast_fp16")];
int32 var_1053 = const()[name = string("op_1053"), val = int32(-1)];
bool input_47_interleave_0 = const()[name = string("input_47_interleave_0"), val = bool(false)];
tensor<fp16, [1, 1, 2048]> input_47_cast_fp16 = concat(axis = var_1053, interleave = input_47_interleave_0, values = (x_17_cast_fp16, var_1051_cast_fp16))[name = string("input_47_cast_fp16")];
tensor<int32, [1]> normed_23_axes_0 = const()[name = string("normed_23_axes_0"), val = tensor<int32, [1]>([-1])];
fp16 var_1059_to_fp16 = const()[name = string("op_1059_to_fp16"), val = fp16(0x1.5p-17)];
tensor<fp16, [1, 1, 2048]> normed_23_cast_fp16 = layer_norm(axes = normed_23_axes_0, epsilon = var_1059_to_fp16, x = input_47_cast_fp16)[name = string("normed_23_cast_fp16")];
tensor<int32, [2]> var_1062_split_sizes_0 = const()[name = string("op_1062_split_sizes_0"), val = tensor<int32, [2]>([1024, 1024])];
int32 var_1062_axis_0 = const()[name = string("op_1062_axis_0"), val = int32(-1)];
tensor<fp16, [1, 1, 1024]> var_1062_cast_fp16_0, tensor<fp16, [1, 1, 1024]> var_1062_cast_fp16_1 = split(axis = var_1062_axis_0, split_sizes = var_1062_split_sizes_0, x = normed_23_cast_fp16)[name = string("op_1062_cast_fp16")];
tensor<fp16, [1024]> layers_3_operator_norm_weight_promoted_to_fp16 = const()[name = string("layers_3_operator_norm_weight_promoted_to_fp16"), val = tensor<fp16, [1024]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(53940800)))];
tensor<fp16, [1, 1, 1024]> hidden_states_cast_fp16 = mul(x = var_1062_cast_fp16_0, y = layers_3_operator_norm_weight_promoted_to_fp16)[name = string("hidden_states_cast_fp16")];
tensor<int32, [3]> var_1068 = const()[name = string("op_1068"), val = tensor<int32, [3]>([0, 2, 1])];
tensor<int32, [1]> var_1071_axes_0 = const()[name = string("op_1071_axes_0"), val = tensor<int32, [1]>([2])];
tensor<fp16, [1, 1024, 1]> var_1069_cast_fp16 = transpose(perm = var_1068, x = hidden_states_cast_fp16)[name = string("transpose_3")];
tensor<fp16, [1, 1024, 1, 1]> var_1071_cast_fp16 = expand_dims(axes = var_1071_axes_0, x = var_1069_cast_fp16)[name = string("op_1071_cast_fp16")];
string BCx_pad_type_0 = const()[name = string("BCx_pad_type_0"), val = string("valid")];
tensor<int32, [2]> BCx_strides_0 = const()[name = string("BCx_strides_0"), val = tensor<int32, [2]>([1, 1])];
tensor<int32, [4]> BCx_pad_0 = const()[name = string("BCx_pad_0"), val = tensor<int32, [4]>([0, 0, 0, 0])];
tensor<int32, [2]> BCx_dilations_0 = const()[name = string("BCx_dilations_0"), val = tensor<int32, [2]>([1, 1])];
int32 BCx_groups_0 = const()[name = string("BCx_groups_0"), val = int32(1)];
tensor<fp16, [1, 3072, 1, 1]> BCx = conv(dilations = BCx_dilations_0, groups = BCx_groups_0, pad = BCx_pad_0, pad_type = BCx_pad_type_0, strides = BCx_strides_0, weight = layers_3_conv_in_proj_weight_palettized, x = var_1071_cast_fp16)[name = string("BCx")];
tensor<int32, [3]> var_1088_split_sizes_0 = const()[name = string("op_1088_split_sizes_0"), val = tensor<int32, [3]>([1024, 1024, 1024])];
int32 var_1088_axis_0 = const()[name = string("op_1088_axis_0"), val = int32(1)];
tensor<fp16, [1, 1024, 1, 1]> var_1088_0, tensor<fp16, [1, 1024, 1, 1]> var_1088_1, tensor<fp16, [1, 1024, 1, 1]> var_1088_2 = split(axis = var_1088_axis_0, split_sizes = var_1088_split_sizes_0, x = BCx)[name = string("op_1088")];
tensor<fp16, [1, 1024, 1, 1]> Bx = mul(x = var_1088_0, y = var_1088_2)[name = string("Bx")];
tensor<int32, [3]> var_1094_begin_0 = const()[name = string("op_1094_begin_0"), val = tensor<int32, [3]>([2, 0, 0])];
tensor<int32, [3]> var_1094_end_0 = const()[name = string("op_1094_end_0"), val = tensor<int32, [3]>([3, 1024, 3])];
tensor<bool, [3]> var_1094_end_mask_0 = const()[name = string("op_1094_end_mask_0"), val = tensor<bool, [3]>([false, true, true])];
tensor<bool, [3]> var_1094_squeeze_mask_0 = const()[name = string("op_1094_squeeze_mask_0"), val = tensor<bool, [3]>([true, false, false])];
tensor<fp16, [1024, 3]> var_1094_cast_fp16 = slice_by_index(begin = var_1094_begin_0, end = var_1094_end_0, end_mask = var_1094_end_mask_0, squeeze_mask = var_1094_squeeze_mask_0, x = conv_state_in)[name = string("op_1094_cast_fp16")];
tensor<int32, [1]> var_1096_axes_0 = const()[name = string("op_1096_axes_0"), val = tensor<int32, [1]>([0])];
tensor<fp16, [1, 1024, 3]> var_1096_cast_fp16 = expand_dims(axes = var_1096_axes_0, x = var_1094_cast_fp16)[name = string("op_1096_cast_fp16")];
tensor<int32, [1]> slot_axes_0 = const()[name = string("slot_axes_0"), val = tensor<int32, [1]>([2])];
tensor<fp16, [1, 1024, 1, 3]> slot_cast_fp16 = expand_dims(axes = slot_axes_0, x = var_1096_cast_fp16)[name = string("slot_cast_fp16")];
tensor<int32, [4]> live_tail_begin_0 = const()[name = string("live_tail_begin_0"), val = tensor<int32, [4]>([0, 0, 0, 1])];
tensor<int32, [4]> live_tail_end_0 = const()[name = string("live_tail_end_0"), val = tensor<int32, [4]>([1, 1024, 1, 1])];
tensor<bool, [4]> live_tail_end_mask_0 = const()[name = string("live_tail_end_mask_0"), val = tensor<bool, [4]>([true, true, true, true])];
tensor<fp16, [1, 1024, 1, 2]> live_tail_cast_fp16 = slice_by_index(begin = live_tail_begin_0, end = live_tail_end_0, end_mask = live_tail_end_mask_0, x = slot_cast_fp16)[name = string("live_tail_cast_fp16")];
int32 var_1105 = const()[name = string("op_1105"), val = int32(-1)];
bool new_state_interleave_0 = const()[name = string("new_state_interleave_0"), val = bool(false)];
tensor<fp16, [1, 1024, 1, 3]> new_state_cast_fp16 = concat(axis = var_1105, interleave = new_state_interleave_0, values = (live_tail_cast_fp16, Bx))[name = string("new_state_cast_fp16")];
tensor<int32, [1]> var_1108_axes_0 = const()[name = string("op_1108_axes_0"), val = tensor<int32, [1]>([0])];
tensor<fp16, [1024, 1, 3]> var_1108_cast_fp16 = squeeze(axes = var_1108_axes_0, x = new_state_cast_fp16)[name = string("op_1108_cast_fp16")];
tensor<int32, [1]> new_slot_axes_0 = const()[name = string("new_slot_axes_0"), val = tensor<int32, [1]>([1])];
tensor<fp16, [1024, 3]> new_slot_cast_fp16 = squeeze(axes = new_slot_axes_0, x = var_1108_cast_fp16)[name = string("new_slot_cast_fp16")];
string conv_out_pad_type_0 = const()[name = string("conv_out_pad_type_0"), val = string("valid")];
int32 conv_out_groups_0 = const()[name = string("conv_out_groups_0"), val = int32(1024)];
tensor<int32, [2]> conv_out_strides_0 = const()[name = string("conv_out_strides_0"), val = tensor<int32, [2]>([1, 1])];
tensor<int32, [4]> conv_out_pad_0 = const()[name = string("conv_out_pad_0"), val = tensor<int32, [4]>([0, 0, 0, 0])];
tensor<int32, [2]> conv_out_dilations_0 = const()[name = string("conv_out_dilations_0"), val = tensor<int32, [2]>([1, 1])];
tensor<fp16, [1024, 1, 1, 3]> layers_3_conv_conv_weight_promoted_to_fp16_palettized = constexpr_lut_to_dense(indices = tensor<uint6, [1024, 1, 1, 3]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(53942912))), lut = tensor<fp16, [32, 1, 1, 1, 64, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(53945280))))[name = string("layers_3_conv_conv_weight_promoted_to_fp16_palettized")];
tensor<fp16, [1, 1024, 1, 1]> conv_out_cast_fp16 = conv(dilations = conv_out_dilations_0, groups = conv_out_groups_0, pad = conv_out_pad_0, pad_type = conv_out_pad_type_0, strides = conv_out_strides_0, weight = layers_3_conv_conv_weight_promoted_to_fp16_palettized, x = new_state_cast_fp16)[name = string("conv_out_cast_fp16")];
tensor<fp16, [1, 1024, 1, 1]> input_51_cast_fp16 = mul(x = var_1088_1, y = conv_out_cast_fp16)[name = string("input_51_cast_fp16")];
string y_pad_type_0 = const()[name = string("y_pad_type_0"), val = string("valid")];
tensor<int32, [2]> y_strides_0 = const()[name = string("y_strides_0"), val = tensor<int32, [2]>([1, 1])];
tensor<int32, [4]> y_pad_0 = const()[name = string("y_pad_0"), val = tensor<int32, [4]>([0, 0, 0, 0])];
tensor<int32, [2]> y_dilations_0 = const()[name = string("y_dilations_0"), val = tensor<int32, [2]>([1, 1])];
int32 y_groups_0 = const()[name = string("y_groups_0"), val = int32(1)];
tensor<fp16, [1024, 1024, 1, 1]> layers_3_conv_out_proj_weight_promoted_to_fp16_palettized = constexpr_lut_to_dense(indices = tensor<uint6, [1024, 1024, 1, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(53949440))), lut = tensor<fp16, [32, 1, 1, 1, 64, 1]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(54735936))))[name = string("layers_3_conv_out_proj_weight_promoted_to_fp16_palettized")];
tensor<fp16, [1, 1024, 1, 1]> y_cast_fp16 = conv(dilations = y_dilations_0, groups = y_groups_0, pad = y_pad_0, pad_type = y_pad_type_0, strides = y_strides_0, weight = layers_3_conv_out_proj_weight_promoted_to_fp16_palettized, x = input_51_cast_fp16)[name = string("y_cast_fp16")];
tensor<int32, [1]> var_1136_axes_0 = const()[name = string("op_1136_axes_0"), val = tensor<int32, [1]>([2])];
tensor<fp16, [1, 1024, 1]> var_1136_cast_fp16 = squeeze(axes = var_1136_axes_0, x = y_cast_fp16)[name = string("op_1136_cast_fp16")];
tensor<int32, [3]> var_1140 = const()[name = string("op_1140"), val = tensor<int32, [3]>([0, 2, 1])];
tensor<fp16, [1, 1, 1024]> op_out_cast_fp16 = transpose(perm = var_1140, x = var_1136_cast_fp16)[name = string("transpose_2")];
tensor<fp16, [1, 1, 1024]> x_cast_fp16 = add(x = x_17_cast_fp16, y = op_out_cast_fp16)[name = string("x_cast_fp16")];
fp16 const_17_promoted_to_fp16 = const()[name = string("const_17_promoted_to_fp16"), val = fp16(-0x1p+0)];
tensor<fp16, [1, 1, 1024]> var_1144_cast_fp16 = mul(x = x_cast_fp16, y = const_17_promoted_to_fp16)[name = string("op_1144_cast_fp16")];
int32 var_1146 = const()[name = string("op_1146"), val = int32(-1)];
bool input_53_interleave_0 = const()[name = string("input_53_interleave_0"), val = bool(false)];
tensor<fp16, [1, 1, 2048]> input_53_cast_fp16 = concat(axis = var_1146, interleave = input_53_interleave_0, values = (x_cast_fp16, var_1144_cast_fp16))[name = string("input_53_cast_fp16")];
tensor<int32, [1]> normed_25_axes_0 = const()[name = string("normed_25_axes_0"), val = tensor<int32, [1]>([-1])];
fp16 var_1152_to_fp16 = const()[name = string("op_1152_to_fp16"), val = fp16(0x1.5p-17)];
tensor<fp16, [1, 1, 2048]> normed_25_cast_fp16 = layer_norm(axes = normed_25_axes_0, epsilon = var_1152_to_fp16, x = input_53_cast_fp16)[name = string("normed_25_cast_fp16")];
tensor<int32, [2]> var_1155_split_sizes_0 = const()[name = string("op_1155_split_sizes_0"), val = tensor<int32, [2]>([1024, 1024])];
int32 var_1155_axis_0 = const()[name = string("op_1155_axis_0"), val = int32(-1)];
tensor<fp16, [1, 1, 1024]> var_1155_cast_fp16_0, tensor<fp16, [1, 1, 1024]> var_1155_cast_fp16_1 = split(axis = var_1155_axis_0, split_sizes = var_1155_split_sizes_0, x = normed_25_cast_fp16)[name = string("op_1155_cast_fp16")];
tensor<fp16, [1024]> layers_3_ffn_norm_weight_promoted_to_fp16 = const()[name = string("layers_3_ffn_norm_weight_promoted_to_fp16"), val = tensor<fp16, [1024]>(BLOBFILE(path = string("@model_path/weights/weight.bin"), offset = uint64(54740096)))];
tensor<fp16, [1, 1, 1024]> normed_cast_fp16 = mul(x = var_1155_cast_fp16_0, y = layers_3_ffn_norm_weight_promoted_to_fp16)[name = string("normed_cast_fp16")];
tensor<int32, [3]> var_1161 = const()[name = string("op_1161"), val = tensor<int32, [3]>([0, 2, 1])];
tensor<int32, [1]> var_1164_axes_0 = const()[name = string("op_1164_axes_0"), val = tensor<int32, [1]>([2])];
tensor<fp16, [1, 1024, 1]> var_1162_cast_fp16 = transpose(perm = var_1161, x = normed_cast_fp16)[name = string("transpose_1")];
tensor<fp16, [1, 1024, 1, 1]> var_1164_cast_fp16 = expand_dims(axes = var_1164_axes_0, x = var_1162_cast_fp16)[name = string("op_1164_cast_fp16")];
string input_57_pad_type_0 = const()[name = string("input_57_pad_type_0"), val = string("valid")];
tensor<int32, [2]> input_57_strides_0 = const()[name = string("input_57_strides_0"), val = tensor<int32, [2]>([1, 1])];
tensor<int32, [4]> input_57_pad_0 = const()[name = string("input_57_pad_0"), val = tensor<int32, [4]>([0, 0, 0, 0])];
tensor<int32, [2]> input_57_dilations_0 = const()[name = string("input_57_dilations_0"), val = tensor<int32, [2]>([1, 1])];
int32 input_57_groups_0 = const()[name = string("input_57_groups_0"), val = int32(1)];
tensor<fp16, [1, 4608, 1, 1]> input_57 = conv(dilations = input_57_dilations_0, groups = input_57_groups_0, pad = input_57_pad_0, pad_type = input_57_pad_type_0, strides = input_57_strides_0, weight = layers_3_feed_forward_w1_weight_palettized, x = var_1164_cast_fp16)[name = string("input_57")];
string b_pad_type_0 = const()[name = string("b_pad_type_0"), val = string("valid")];
tensor<int32, [2]> b_strides_0 = const()[name = string("b_strides_0"), val = tensor<int32, [2]>([1, 1])];
tensor<int32, [4]> b_pad_0 = const()[name = string("b_pad_0"), val = tensor<int32, [4]>([0, 0, 0, 0])];
tensor<int32, [2]> b_dilations_0 = const()[name = string("b_dilations_0"), val = tensor<int32, [2]>([1, 1])];
int32 b_groups_0 = const()[name = string("b_groups_0"), val = int32(1)];
tensor<fp16, [1, 4608, 1, 1]> b = conv(dilations = b_dilations_0, groups = b_groups_0, pad = b_pad_0, pad_type = b_pad_type_0, strides = b_strides_0, weight = layers_3_feed_forward_w3_weight_palettized, x = var_1164_cast_fp16)[name = string("b")];
tensor<fp16, [1, 4608, 1, 1]> var_1192 = silu(x = input_57)[name = string("op_1192")];
tensor<fp16, [1, 4608, 1, 1]> input = mul(x = var_1192, y = b)[name = string("input")];
string mlp_13_pad_type_0 = const()[name = string("mlp_13_pad_type_0"), val = string("valid")];
tensor<int32, [2]> mlp_13_strides_0 = const()[name = string("mlp_13_strides_0"), val = tensor<int32, [2]>([1, 1])];
tensor<int32, [4]> mlp_13_pad_0 = const()[name = string("mlp_13_pad_0"), val = tensor<int32, [4]>([0, 0, 0, 0])];
tensor<int32, [2]> mlp_13_dilations_0 = const()[name = string("mlp_13_dilations_0"), val = tensor<int32, [2]>([1, 1])];
int32 mlp_13_groups_0 = const()[name = string("mlp_13_groups_0"), val = int32(1)];
tensor<fp16, [1, 1024, 1, 1]> mlp_13 = conv(dilations = mlp_13_dilations_0, groups = mlp_13_groups_0, pad = mlp_13_pad_0, pad_type = mlp_13_pad_type_0, strides = mlp_13_strides_0, weight = layers_3_feed_forward_w2_weight_palettized, x = input)[name = string("mlp_13")];
tensor<int32, [1]> var_1206_axes_0 = const()[name = string("op_1206_axes_0"), val = tensor<int32, [1]>([2])];
tensor<fp16, [1, 1024, 1]> var_1206 = squeeze(axes = var_1206_axes_0, x = mlp_13)[name = string("op_1206")];
tensor<int32, [3]> var_1210 = const()[name = string("op_1210"), val = tensor<int32, [3]>([0, 2, 1])];
tensor<fp16, [1, 1, 1024]> mlp = transpose(perm = var_1210, x = var_1206)[name = string("transpose_0")];
tensor<fp16, [1, 1, 1024]> hidden_out = add(x = x_cast_fp16, y = mlp)[name = string("op_1213_cast_fp16")];
int32 var_1216_axis_0 = const()[name = string("op_1216_axis_0"), val = int32(0)];
tensor<fp16, [3, 1024, 3]> conv_state_out = stack(axis = var_1216_axis_0, values = (var_232_cast_fp16, var_947_cast_fp16, new_slot_cast_fp16))[name = string("op_1216_cast_fp16")];
tensor<int32, [1]> var_1219_axes_0 = const()[name = string("op_1219_axes_0"), val = tensor<int32, [1]>([0])];
tensor<fp16, [1, 1, 512, 1, 1]> var_1219 = expand_dims(axes = var_1219_axes_0, x = k_slice)[name = string("op_1219")];
tensor<int32, [1]> var_1222_axes_0 = const()[name = string("op_1222_axes_0"), val = tensor<int32, [1]>([0])];
tensor<fp16, [1, 1, 512, 1, 1]> var_1222 = expand_dims(axes = var_1222_axes_0, x = var_428)[name = string("op_1222")];
int32 var_1224 = const()[name = string("op_1224"), val = int32(0)];
bool var_1225_interleave_0 = const()[name = string("op_1225_interleave_0"), val = bool(false)];
tensor<fp16, [2, 1, 512, 1, 1]> kv_slice_out = concat(axis = var_1224, interleave = var_1225_interleave_0, values = (var_1219, var_1222))[name = string("op_1225")];
tensor<fp16, [1, 1, 1024, 1]> update_mask_tmp = identity(x = update_mask)[name = string("update_mask_tmp")];
} -> (hidden_out, kv_slice_out, conv_state_out);
}