Skip to content

Commit d1c5b06

Browse files
committed
Adapt to the new seq_len cli-flag of gemma.cpp
1 parent 8083750 commit d1c5b06

10 files changed

Lines changed: 41 additions & 14 deletions

File tree

README.md

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -171,6 +171,7 @@ Available options and default values:
171171

172172
```lua
173173
{
174+
seq_len = 2048, -- Sequence length, capped by max context window of model.
174175
max_generated_tokens = 2048, -- Maximum number of tokens to generate.
175176
prefill_tbatch = 256, -- Prefill: max tokens per batch.
176177
decode_qbatch = 16, -- Decode: max queries per batch.

demo/src/init.lua

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,7 @@ function config()
99
weights = "4b-it-sfp.sbs"
1010
},
1111
session = {
12+
seq_len = 8192,
1213
temperature = 0.4,
1314
top_k = 5
1415
},

demo/src/init_fn.lua

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,7 @@ function config()
99
weights = "12b-it-sfp.sbs"
1010
},
1111
session = {
12+
seq_len = 8192,
1213
temperature = 0.4,
1314
top_k = 5
1415
},

demo/src/init_huggingface.lua

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@ function config()
1010
weights = "4b-it-sfp.sbs"
1111
},
1212
session = {
13+
seq_len = 8192,
1314
temperature = 0.4,
1415
top_k = 5
1516
},

demo/src/init_kaggle.lua

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@ function config()
66
weights = "4b-it-sfp.sbs"
77
},
88
session = {
9+
seq_len = 8192,
910
prefill_tbatch = 64,
1011
temperature = 0.4,
1112
top_k = 5

demo/src/init_vlm.lua

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,7 @@ function config()
1111
weights = "4b-it-sfp.sbs"
1212
},
1313
session = {
14+
seq_len = 8192,
1415
temperature = 0.4,
1516
top_k = 5
1617
},

examples/normal_mode.lua

Lines changed: 10 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,10 @@ if args.help then
55
"resty normal_mode.lua [options]"
66
)
77
print(" --image: Path of image file (PPM format: P6, binary).")
8+
print(" --seq_len: Sequence length, capped by max context window of model. (default: 4096)")
9+
print(" --max_generated_tokens: Maximum number of tokens to generate. (default: 2048)")
10+
print(" --temperature: Temperature for top-K. (default: 1.0)")
11+
print(" --top_k: Number of top-K tokens to sample from. (default: 5)")
812
print(" --kv_cache: Path of KV cache file.")
913
print(" --stats: Print statistics at end of turn.")
1014
return
@@ -29,7 +33,12 @@ if args.image then
2933
end
3034

3135
-- Create a chat session
32-
local session, err = gemma:session({top_k = 5})
36+
local session, err = gemma:session({
37+
seq_len = tonumber(args.seq_len) or 4096,
38+
max_generated_tokens = tonumber(args.max_generated_tokens),
39+
temperature = tonumber(args.temperature),
40+
top_k = tonumber(args.top_k) or 5
41+
})
3342
if not session then
3443
error("Opoos! "..err)
3544
end

examples/stream_mode.lua

Lines changed: 10 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,10 @@ if args.help then
55
"resty stream_mode.lua [options]"
66
)
77
print(" --image: Path of image file (PPM format: P6, binary).")
8+
print(" --seq_len: Sequence length, capped by max context window of model. (default: 4096)")
9+
print(" --max_generated_tokens: Maximum number of tokens to generate. (default: 2048)")
10+
print(" --temperature: Temperature for top-K. (default: 1.0)")
11+
print(" --top_k: Number of top-K tokens to sample from. (default: 5)")
812
print(" --kv_cache: Path of KV cache file.")
913
print(" --stats: Print statistics at end of turn.")
1014
return
@@ -29,7 +33,12 @@ if args.image then
2933
end
3034

3135
-- Create a chat session
32-
local session, err = gemma:session({top_k = 5})
36+
local session, err = gemma:session({
37+
seq_len = tonumber(args.seq_len) or 4096,
38+
max_generated_tokens = tonumber(args.max_generated_tokens),
39+
temperature = tonumber(args.temperature),
40+
top_k = tonumber(args.top_k) or 5
41+
})
3342
if not session then
3443
error("Opoos! "..err)
3544
end

src/session.cpp

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -383,6 +383,7 @@ int session::create(lua_State* L) {
383383
auto nargs = lua_gettop(L);
384384
auto inst = instance::check(L, 1);
385385
constexpr const char* available_options[] = {
386+
"--seq_len",
386387
"--max_generated_tokens",
387388
"--prefill_tbatch",
388389
"--decode_qbatch",

tools/dump_prompt.lua

Lines changed: 14 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,7 @@ if args.help then
2020
print(" --tokenizer: Path of tokenizer model file. (default: tokenizer.spm)")
2121
print(" --weights: Path of model weights file. (default: 4b-it-sfp.sbs)")
2222
print(" --map: Enable memory-mapping? -1 = auto, 0 = no, 1 = yes. (default: -1)")
23+
print(" --seq_len: Sequence length, capped by max context window of model. (default: 4096)")
2324
print(" --max_generated_tokens: Maximum number of tokens to generate. (default: 2048)")
2425
print(" --prefill_tbatch: Maximum batch size during prefill phase (default: 256)")
2526
print(" --temperature: Temperature for top-K. (default: 1.0)")
@@ -43,14 +44,14 @@ end
4344

4445
-- Config global scheduler
4546
local ok, err = require("cgemma").scheduler.config({
46-
num_threads = args.num_threads,
47-
pin = args.pin,
48-
skip_packages = args.skip_packages,
49-
max_packages = args.max_packages,
50-
skip_clusters = args.skip_clusters,
51-
max_clusters = args.max_clusters,
52-
skip_lps = args.skip_lps,
53-
max_lps = args.max_lps
47+
num_threads = tonumber(args.num_threads),
48+
pin = tonumber(args.pin),
49+
skip_packages = tonumber(args.skip_packages),
50+
max_packages = tonumber(args.max_packages),
51+
skip_clusters = tonumber(args.skip_clusters),
52+
max_clusters = tonumber(args.max_clusters),
53+
skip_lps = tonumber(args.skip_lps),
54+
max_lps = tonumber(args.max_lps)
5455
})
5556
if not ok then
5657
error("Opoos! "..err)
@@ -78,10 +79,11 @@ end
7879

7980
-- Create a session
8081
local session, err = gemma:session({
81-
max_generated_tokens = args.max_generated_tokens,
82-
prefill_tbatch = args.prefill_tbatch,
83-
temperature = args.temperature,
84-
top_k = args.top_k,
82+
seq_len = tonumber(args.seq_len) or 4096,
83+
max_generated_tokens = tonumber(args.max_generated_tokens),
84+
prefill_tbatch = tonumber(args.prefill_tbatch),
85+
temperature = tonumber(args.temperature),
86+
top_k = tonumber(args.top_k),
8587
no_wrapping = true
8688
})
8789
if not session then

0 commit comments

Comments
 (0)