TODO: maybe separate the inner implementation into a separate function like with the non-sliding window equivalent once sliding-window hybrid caches are a thing.
| 2099 | // like with the non-sliding window equivalent |
| 2100 | // once sliding-window hybrid caches are a thing. |
| 2101 | llm_graph_input_attn_kv_iswa * llm_graph_context::build_attn_inp_kv_iswa() const { |
| 2102 | const auto * mctx_cur = static_cast<const llama_kv_cache_iswa_context *>(mctx); |
| 2103 | |
| 2104 | auto inp = std::make_unique<llm_graph_input_attn_kv_iswa>(hparams, cparams, mctx_cur); |
| 2105 | |
| 2106 | const auto n_stream = cparams.kv_unified ? 1 : ubatch.n_seqs_unq; |
| 2107 | |
| 2108 | { |
| 2109 | const auto n_kv = mctx_cur->get_base()->get_n_kv(); |
| 2110 | |
| 2111 | inp->self_k_idxs = mctx_cur->get_base()->build_input_k_idxs(ctx0, ubatch); |
| 2112 | inp->self_v_idxs = mctx_cur->get_base()->build_input_v_idxs(ctx0, ubatch); |
| 2113 | |
| 2114 | inp->self_kq_mask = ggml_new_tensor_4d(ctx0, GGML_TYPE_F32, n_kv, n_tokens/n_stream, 1, n_stream); |
| 2115 | ggml_set_input(inp->self_kq_mask); |
| 2116 | ggml_set_name(inp->self_kq_mask, "self_kq_mask"); |
| 2117 | |
| 2118 | inp->self_kq_mask_cnv = cparams.flash_attn ? ggml_cast(ctx0, inp->self_kq_mask, GGML_TYPE_F16) : inp->self_kq_mask; |
| 2119 | ggml_set_name(inp->self_kq_mask_cnv, "self_kq_mask_cnv"); |
| 2120 | } |
| 2121 | |
| 2122 | { |
| 2123 | GGML_ASSERT(hparams.swa_type != LLAMA_SWA_TYPE_NONE && "Use llama_kv_cache for non-SWA"); |
| 2124 | |
| 2125 | const auto n_kv = mctx_cur->get_swa()->get_n_kv(); |
| 2126 | |
| 2127 | inp->self_k_idxs_swa = mctx_cur->get_swa()->build_input_k_idxs(ctx0, ubatch); |
| 2128 | inp->self_v_idxs_swa = mctx_cur->get_swa()->build_input_v_idxs(ctx0, ubatch); |
| 2129 | |
| 2130 | inp->self_kq_mask_swa = ggml_new_tensor_4d(ctx0, GGML_TYPE_F32, n_kv, n_tokens/n_stream, 1, n_stream); |
| 2131 | ggml_set_input(inp->self_kq_mask_swa); |
| 2132 | ggml_set_name(inp->self_kq_mask_swa, "self_kq_mask_swa"); |
| 2133 | |
| 2134 | inp->self_kq_mask_swa_cnv = cparams.flash_attn ? ggml_cast(ctx0, inp->self_kq_mask_swa, GGML_TYPE_F16) : inp->self_kq_mask_swa; |
| 2135 | ggml_set_name(inp->self_kq_mask_swa_cnv, "self_kq_mask_swa_cnv"); |
| 2136 | } |
| 2137 | |
| 2138 | return (llm_graph_input_attn_kv_iswa *) res->add_input(std::move(inp)); |
| 2139 | } |
| 2140 | |
| 2141 | ggml_tensor * llm_graph_context::build_rs( |
| 2142 | ggml_tensor * s, |
nothing calls this directly
no test coverage detected