}
}
-std::vector<common_sampler_type> common_sampler_types_from_names(const std::vector<std::string> & names, bool allow_alt_names) {
- std::unordered_map<std::string, common_sampler_type> sampler_canonical_name_map {
- { "dry", COMMON_SAMPLER_TYPE_DRY },
- { "top_k", COMMON_SAMPLER_TYPE_TOP_K },
- { "top_p", COMMON_SAMPLER_TYPE_TOP_P },
- { "top_n_sigma", COMMON_SAMPLER_TYPE_TOP_N_SIGMA },
- { "typ_p", COMMON_SAMPLER_TYPE_TYPICAL_P },
- { "min_p", COMMON_SAMPLER_TYPE_MIN_P },
- { "temperature", COMMON_SAMPLER_TYPE_TEMPERATURE },
- { "xtc", COMMON_SAMPLER_TYPE_XTC },
- { "infill", COMMON_SAMPLER_TYPE_INFILL },
- { "penalties", COMMON_SAMPLER_TYPE_PENALTIES },
- { "adaptive_p", COMMON_SAMPLER_TYPE_ADAPTIVE_P },
- };
-
- // since samplers names are written multiple ways
- // make it ready for both system names and input names
- std::unordered_map<std::string, common_sampler_type> sampler_alt_name_map {
- { "top-k", COMMON_SAMPLER_TYPE_TOP_K },
- { "top-p", COMMON_SAMPLER_TYPE_TOP_P },
- { "top-n-sigma", COMMON_SAMPLER_TYPE_TOP_N_SIGMA },
- { "nucleus", COMMON_SAMPLER_TYPE_TOP_P },
- { "typical-p", COMMON_SAMPLER_TYPE_TYPICAL_P },
- { "typical", COMMON_SAMPLER_TYPE_TYPICAL_P },
- { "typ-p", COMMON_SAMPLER_TYPE_TYPICAL_P },
- { "typ", COMMON_SAMPLER_TYPE_TYPICAL_P },
- { "min-p", COMMON_SAMPLER_TYPE_MIN_P },
- { "temp", COMMON_SAMPLER_TYPE_TEMPERATURE },
- { "adaptive-p", COMMON_SAMPLER_TYPE_ADAPTIVE_P },
- };
+std::vector<common_sampler_type> common_sampler_types_from_names(const std::vector<std::string> & names) {
+ // sampler names can be written multiple ways; generate aliases from canonical names
+ static const auto sampler_name_map = []{
+ // canonical sampler name mapping
+ std::unordered_map<std::string, common_sampler_type> canonical_name_map {
+ { "dry", COMMON_SAMPLER_TYPE_DRY },
+ { "top_k", COMMON_SAMPLER_TYPE_TOP_K },
+ { "top_p", COMMON_SAMPLER_TYPE_TOP_P },
+ { "top_n_sigma", COMMON_SAMPLER_TYPE_TOP_N_SIGMA },
+ { "typ_p", COMMON_SAMPLER_TYPE_TYPICAL_P },
+ { "min_p", COMMON_SAMPLER_TYPE_MIN_P },
+ { "temperature", COMMON_SAMPLER_TYPE_TEMPERATURE },
+ { "xtc", COMMON_SAMPLER_TYPE_XTC },
+ { "infill", COMMON_SAMPLER_TYPE_INFILL },
+ { "penalties", COMMON_SAMPLER_TYPE_PENALTIES },
+ { "adaptive_p", COMMON_SAMPLER_TYPE_ADAPTIVE_P }
+ };
+ std::unordered_map<std::string, common_sampler_type> alias_name_map;
+ for (const auto & entry : canonical_name_map) {
+ const std::string & canonical = entry.first;
+ if (canonical.find('_') == std::string::npos) {
+ continue;
+ }
+ // kebab-case: "top-k", "min-p", etc.
+ {
+ std::string kebab_case = canonical;
+ std::replace(kebab_case.begin(), kebab_case.end(), '_', '-');
+ alias_name_map.insert({kebab_case, entry.second});
+ }
+ // no dash: "topk", "minp", etc.
+ {
+ std::string no_dash = canonical;
+ no_dash.erase(std::remove(no_dash.begin(), no_dash.end(), '_'), no_dash.end());
+ alias_name_map.insert({no_dash, entry.second});
+ }
+ }
+ // misc. aliases
+ alias_name_map.insert({"nucleus", COMMON_SAMPLER_TYPE_TOP_P});
+ alias_name_map.insert({"temp", COMMON_SAMPLER_TYPE_TEMPERATURE});
+ alias_name_map.insert({"typ", COMMON_SAMPLER_TYPE_TYPICAL_P});
+ // include aliases + canonical names in the complete mapping
+ alias_name_map.merge(canonical_name_map);
+ return alias_name_map;
+ }();
std::vector<common_sampler_type> samplers;
samplers.reserve(names.size());
for (const auto & name : names) {
- auto sampler = sampler_canonical_name_map.find(name);
- if (sampler != sampler_canonical_name_map.end()) {
+ std::string name_lower = name;
+ std::transform(name_lower.begin(), name_lower.end(), name_lower.begin(), ::tolower);
+ auto sampler = sampler_name_map.find(name_lower);
+ if (sampler != sampler_name_map.end()) {
samplers.push_back(sampler->second);
continue;
}
- if (allow_alt_names) {
- sampler = sampler_alt_name_map.find(name);
- if (sampler != sampler_alt_name_map.end()) {
- samplers.push_back(sampler->second);
- continue;
- }
- }
- LOG_WRN("%s: unable to match sampler by name '%s'\n", __func__, name.c_str());
+ LOG_WRN("%s: unable to match sampler by name '%s'\n", __func__, name_lower.c_str());
}
return samplers;