| 20 | } |
| 21 | |
| 22 | void Config::from_yalm(YALMData& yalm, int context) { |
| 23 | dim = std::stoi(yalm.metadata.at("dim").get<std::string>()); |
| 24 | hidden_dim = std::stoi(yalm.metadata.at("hidden_dim").get<std::string>()); |
| 25 | n_layers = std::stoi(yalm.metadata.at("n_layers").get<std::string>()); |
| 26 | n_heads = std::stoi(yalm.metadata.at("n_heads").get<std::string>()); |
| 27 | vocab_size = std::stoi(yalm.metadata.at("vocab_size").get<std::string>()); |
| 28 | // mixture of experts |
| 29 | n_shared_experts = yalm.metadata.contains("n_shared_experts") ? std::stoi(yalm.metadata.at("n_shared_experts").get<std::string>()) : 0; |
| 30 | n_routed_experts = yalm.metadata.contains("n_routed_experts") ? std::stoi(yalm.metadata.at("n_routed_experts").get<std::string>()) : 0; |
| 31 | n_active_routed = yalm.metadata.contains("n_active_routed") ? std::stoi(yalm.metadata.at("n_active_routed").get<std::string>()) : 0; |
| 32 | moe_intermediate_size = yalm.metadata.contains("moe_intermediate_size") ? std::stoi(yalm.metadata.at("moe_intermediate_size").get<std::string>()) : 0; |
| 33 | routed_scaling_factor = yalm.metadata.contains("routed_scaling_factor") ? std::stof(yalm.metadata.at("routed_scaling_factor").get<std::string>()) : 1.0; |
| 34 | n_group = yalm.metadata.contains("n_group") ? std::stoi(yalm.metadata.at("n_group").get<std::string>()) : 1; |
| 35 | norm_topk_prob = yalm.metadata.contains("norm_topk_prob") ? yalm.metadata.at("norm_topk_prob").get<std::string>() == "True" : false; |
| 36 | std::string scoring_func_str = yalm.metadata.value("scoring_func", "softmax"); |
| 37 | if (scoring_func_str == "softmax") { |
| 38 | scoring_func = ScoringFunc::SOFTMAX; |
| 39 | } else if (scoring_func_str == "sigmoid") { |
| 40 | scoring_func = ScoringFunc::SIGMOID; |
| 41 | } else { |
| 42 | std::cerr << "unsupported scoring_func '" << scoring_func_str << "', defaulting to softmax" << std::endl; |
| 43 | scoring_func = ScoringFunc::SOFTMAX; |
| 44 | } |
| 45 | topk_group = yalm.metadata.contains("topk_group") ? std::stoi(yalm.metadata.at("topk_group").get<std::string>()) : 0; |
| 46 | std::string topk_method_str = yalm.metadata.value("topk_method", ""); |
| 47 | if (topk_method_str == "greedy") { |
| 48 | topk_method = TopKMethod::GREEDY; |
| 49 | } else if (topk_method_str == "group_limited_greedy") { |
| 50 | topk_method = TopKMethod::GROUP_LIMITED_GREEDY; |
| 51 | } else if (topk_method_str == "noaux_tc") { |
| 52 | topk_method = TopKMethod::NOAUX_TC; |
| 53 | assert(false && "TODO: support for Deepseek v3"); |
| 54 | } else { |
| 55 | std::cerr << "unsupported topk_method '" << topk_method_str << "', defaulting to greedy" << std::endl; |
| 56 | topk_method = TopKMethod::GREEDY; |
| 57 | } |
| 58 | has_moegate_bias = yalm.metadata.at("arch").get<std::string>() == "DeepseekV3ForCausalLM"; |
| 59 | // multi-latent attention |
| 60 | use_mla = yalm.metadata.contains("use_mla") ? |
| 61 | static_cast<bool>(std::stoi(yalm.metadata.at("use_mla").get<std::string>())) : false; |
| 62 | kv_lora_rank = yalm.metadata.contains("kv_lora_rank") ? std::stoi(yalm.metadata.at("kv_lora_rank").get<std::string>()) : 0; |
| 63 | q_lora_rank = yalm.metadata.contains("q_lora_rank") ? std::stoi(yalm.metadata.at("q_lora_rank").get<std::string>()) : 0; |
| 64 | qk_nope_head_dim = yalm.metadata.contains("qk_nope_head_dim") ? std::stoi(yalm.metadata.at("qk_nope_head_dim").get<std::string>()) : 0; |
| 65 | qk_rope_head_dim = yalm.metadata.contains("qk_rope_head_dim") ? std::stoi(yalm.metadata.at("qk_rope_head_dim").get<std::string>()) : 0; |
| 66 | v_head_dim = yalm.metadata.contains("v_head_dim") ? std::stoi(yalm.metadata.at("v_head_dim").get<std::string>()) : 0; |
| 67 | head_dim = qk_nope_head_dim + qk_rope_head_dim; |
| 68 | |
| 69 | max_seq_len = std::stoi(yalm.metadata.at("max_seq_len").get<std::string>()); |
| 70 | if (context) { |
| 71 | max_seq_len = std::min(max_seq_len, context); |
| 72 | } |
| 73 | |
| 74 | rope_theta = std::stof(yalm.metadata.at("rope_theta").get<std::string>()); |
| 75 | norm_eps = std::stof(yalm.metadata.value("norm_eps", "1e-5")); |
| 76 | |
| 77 | std::string act_str = yalm.metadata.value("act_type", "gelu"); |
| 78 | if (act_str == "gelu") { |
| 79 | act = ActivationType::GELU; |