Logo

index : gist

Random things

  • summary
  • about
  • tree
  • log
  • branches
<< path: root/public/gist.git/html/python/ml_grammar.py blob: 88cddefe49e82146bfa49aae47a22658fc085556 [raw] [clear marker]

        
0#!/bin/python
1
2from dataclasses import dataclass
3from shutil import get_terminal_size
4from sys import exit, argv
5from subprocess import run
6from pathlib import Path
7
8
9
10MODEL_PATH = "/drives/drive3/ML/LLama-3-8B-grammar-correction/Llama-3-8B-grammar-correction.Q6_K.gguf"
11TEMPERATURE = 0.1 # Recommendation by model
12TOKENS_TO_PREDICT = 256
13GPU_LAYERS = 99 # Also possible: "auto", "all"
14
15
16LLAMA_TOOL = "llama-completion"
17
18SYSTEM_PROMPT = "Correct the grammar and spelling of the following text, " \
19 "but omit what was already correct."
20
21
22@dataclass
23class Args:
24 fp: string = ""
25 debug: bool = False
26
27
28def main():
29 fp_abs = ""
30 parsed_args = arg_parse()
31
32 if not parsed_args.debug:
33 fp = Path(parsed_args.fp)
34 fp_abs = fp.absolute()
35
36 if not fp.exists():
37 print("Error: File path does not exist:", fp_abs)
38 exit(1)
39
40 print("Loading ...")
41 run_model(parsed_args, fp_abs)
42
43
44def arg_parse():
45 parsed_args = Args()
46 args = argv
47 max_args = 2
48
49 if len(args) == 1 or len(args) > max_args:
50 print("Insufficient arguments. Expecting a path to a text file")
51 exit(1)
52
53 if args[1] == "-debug":
54 parsed_args.debug = True
55 else:
56 parsed_args.fp = args[1]
57
58 return parsed_args
59
60
61def run_model(parsed_args, fp):
62 cmd = [
63 LLAMA_TOOL,
64 "-m", MODEL_PATH,
65 "-ngl", str(GPU_LAYERS),
66 "--single-turn", # Exits llama after one prompt
67 "--temp", str(TEMPERATURE),
68 "--n-predict", str(TOKENS_TO_PREDICT),
69 "--no-display-prompt",
70 "--system-prompt", SYSTEM_PROMPT,
71 "-f", fp,
72 ]
73
74 if parsed_args.debug:
75 print(" ".join(cmd))
76 return
77
78 try:
79 result = run(cmd, capture_output=True, text=True)
80 except FileNotFoundError:
81 print(f"Could not find `{LLAMA_TOOL}` in $PATH")
82 exit(1)
83 except Exception as e:
84 print(e)
85 exit(1)
86
87 output_label = "-- OUTPUT "
88 term_width = get_terminal_size(fallback=(80, 24)).columns
89 width_remainder = abs(term_width - len(output_label))
90
91 if result.stdout:
92 print(f"{output_label}{'-' * width_remainder}")
93 print(result.stdout)
94 if result.returncode != 0 and result.stderr:
95 print("*** *** ERROR *** ***")
96 print(result.stderr)
97
98 exit(result.returncode)
99
100
101if __name__ == "__main__":
102 try:
103 main()
104 except KeyboardInterrupt:
105 print("Terminated by user.")
106 exit(1)
107
108
Copyright 2026  E766CB298A6D1E64 | Git-Thing heavily inspired by cgit