adyoi commited on
Commit
13c9a92
·
verified ·
1 Parent(s): 068c6b9

Delete generate_tf.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. generate_tf.py +0 -68
generate_tf.py DELETED
@@ -1,68 +0,0 @@
1
- import argparse
2
- import os
3
- import random
4
-
5
- import numpy as np
6
- import tensorflow as tf
7
-
8
- from indigo.common import load_meta
9
- from indigo.model_keras import build_gpt, generate
10
- from indigo.tokenizer import CharTokenizer
11
-
12
- CONFIG_KEYS = ("vocab_size", "block_size", "n_layer", "n_head", "n_embd", "dropout")
13
-
14
-
15
- def load_model(path):
16
- from safetensors.numpy import load_file
17
-
18
- meta = load_meta(path)
19
- if meta.get("backend") != "tensorflow":
20
- raise SystemExit(
21
- f"{path} berasal dari backend {meta.get('backend')}, gunakan generate.py (PyTorch)"
22
- )
23
- state = load_file(path)
24
- by_path = {k.replace("/", "_"): v for k, v in state.items()}
25
- model = build_gpt(**{k: meta["config"][k] for k in CONFIG_KEYS})
26
- missing = [v.path for v in model.weights if v.path.replace("/", "_") not in by_path]
27
- if missing:
28
- raise SystemExit(f"bobot tidak cocok dengan checkpoint: {missing[:5]}")
29
- model.set_weights([by_path[v.path.replace("/", "_")] for v in model.weights])
30
- return model, meta
31
-
32
-
33
- def main():
34
- parser = argparse.ArgumentParser(description="Generate teks dari checkpoint Indigo (TensorFlow)")
35
- parser.add_argument("--ckpt", default="out_tf/indigo_best.safetensors")
36
- parser.add_argument("--prompt", default="")
37
- parser.add_argument("--max-new", type=int, default=300)
38
- parser.add_argument("--temperature", type=float, default=0.8)
39
- parser.add_argument("--top-k", type=int, default=40)
40
- parser.add_argument("--seed", type=int, default=None)
41
- parser.add_argument("--device", default="auto", choices=["auto", "cpu", "gpu"])
42
- args = parser.parse_args()
43
-
44
- if args.device == "cpu":
45
- tf.config.set_visible_devices([], "GPU")
46
- if args.seed is not None:
47
- random.seed(args.seed)
48
- np.random.seed(args.seed)
49
- tf.random.set_seed(args.seed)
50
-
51
- model, meta = load_model(args.ckpt)
52
- tokenizer = CharTokenizer(meta["vocab"])
53
-
54
- ids = tokenizer.encode(args.prompt) or [0]
55
- idx = tf.constant([ids], dtype=tf.int64)
56
- out = generate(
57
- model,
58
- idx,
59
- args.max_new,
60
- block_size=meta["config"]["block_size"],
61
- temperature=args.temperature,
62
- top_k=args.top_k,
63
- )
64
- print(tokenizer.decode(out.numpy()[0].tolist()))
65
-
66
-
67
- if __name__ == "__main__":
68
- main()