Repository navigation
Expand file tree
/
Copy pathrunner_utils.h
More file actions
118 lines (107 loc) · 2.95 KB
/
Copy pathrunner_utils.h
File metadata and controls
118 lines (107 loc) · 2.95 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
/*
* Copyright (c) Meta Platforms, Inc. and affiliates.
* All rights reserved.
*
* This source code is licensed under the BSD-style license found in the
* LICENSE file in the root directory of this source tree.
*/
#pragma once
#include <cstdint>
#include <optional>
#include <string>
#include <unordered_set>
#include <executorch/extension/llm/runner/llm_runner_helper.h>
#include <executorch/extension/module/module.h>
#include <executorch/runtime/core/exec_aten/util/scalar_type_util.h>
#include <pytorch/tokenizers/tokenizer.h>
namespace executorch {
namespace backends {
namespace mlx {
namespace examples {
namespace llm {
struct StopTokens {
std::unordered_set<uint64_t> ids;
std::optional<uint64_t> turn_end_id;
};
inline const char* turn_end_piece(const std::string& chat) {
if (chat == "llama3") {
return "<|eot_id|>";
}
if (chat == "gemma") {
return "<end_of_turn>";
}
if (chat == "gemma4") {
return "<turn|>";
}
return nullptr;
}
inline bool resolve_stop_tokens(
tokenizers::Tokenizer& tokenizer,
::executorch::extension::Module& module,
const std::string& chat,
StopTokens& out) {
out.ids = ::executorch::extension::llm::get_eos_ids(&tokenizer, &module);
out.turn_end_id.reset();
if (chat == "0") {
return true;
}
const char* piece = turn_end_piece(chat);
if (piece == nullptr) {
return false;
}
auto id = tokenizer.piece_to_id(piece);
if (!id.ok()) {
return false;
}
out.turn_end_id = *id;
out.ids.insert(*id);
return true;
}
inline bool wrap_turn(
const std::string& chat,
const std::string& prompt,
bool with_bos,
std::string& out) {
if (chat == "0") {
out = prompt;
} else if (chat == "llama3") {
out = std::string(with_bos ? "<|begin_of_text|>" : "") +
"<|start_header_id|>user<|end_header_id|>\n\n" + prompt +
"<|eot_id|><|start_header_id|>assistant<|end_header_id|>\n\n";
} else if (chat == "gemma") {
out = std::string(with_bos ? "<bos>" : "") + "<start_of_turn>user\n" +
prompt + "<end_of_turn>\n<start_of_turn>model\n";
} else if (chat == "gemma4") {
out = std::string(with_bos ? "<bos>" : "") + "<|turn>user\n" + prompt +
"<turn|>\n<|turn>model\n";
} else {
return false;
}
return true;
}
inline int storage_dtype(const std::string& name) {
using ScalarType = ::executorch::runtime::etensor::ScalarType;
if (name == "bf16") {
return static_cast<int>(ScalarType::BFloat16);
}
if (name == "fp16") {
return static_cast<int>(ScalarType::Half);
}
if (name == "fp32") {
return static_cast<int>(ScalarType::Float);
}
return -1;
}
inline int resolve_kv_storage_dtype(
const std::string& override_name,
::executorch::aten::ScalarType activation_dtype) {
if (!override_name.empty()) {
return storage_dtype(override_name);
}
return static_cast<int>(activation_dtype);
}
} // namespace llm
} // namespace examples
} // namespace mlx
} // namespace backends
} // namespace executorch