-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmain.py
More file actions
102 lines (82 loc) · 2.54 KB
/
Copy pathmain.py
File metadata and controls
102 lines (82 loc) · 2.54 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
# cheetah distributed AI
from __future__ import annotations
import argparse
import atexit
import os
import sys
from typing import Optional
from cheetah.tui import main_menu
try:
from dotenv import load_dotenv # type: ignore
except Exception: # pragma: no cover - optional dependency
load_dotenv = None # type: ignore[assignment]
try:
from tinygrad.device import Device
except Exception: # pragma: no cover - tinygrad missing or import failure
Device = None # type: ignore[assignment]
else:
def _cleanup_tinygrad_devices() -> None:
if Device is None:
return
opened = getattr(Device, "_opened_devices", set())
for dev_name in list(opened):
try:
device = Device[dev_name]
except Exception:
continue
try:
device.synchronize()
except Exception:
pass
try:
device.finalize()
except Exception:
pass
atexit.register(_cleanup_tinygrad_devices)
class CheetahApp(main_menu.MainMenu):
"""Main menu app with optional training and chat defaults."""
def __init__(
self,
chat_default: Optional[str] = None,
offline_mode: bool = False
) -> None:
super().__init__(
chat_default=chat_default,
offline_mode=offline_mode
)
def parse_cli_args(argv: list[str]) -> tuple[Optional[str], bool]:
parser = argparse.ArgumentParser(
description="Cheetah TUI launcher",
add_help=True
)
parser.add_argument(
"--chat-model",
dest="chat_model",
type=str,
default=None,
help="Default chat model identifier passed to the chat screen."
)
parser.add_argument(
"--offline-mode",
dest="offline_mode",
action="store_true",
help="Force chat mode to operate offline (use cached models only)."
)
parsed = parser.parse_args(argv)
return parsed.chat_model, parsed.offline_mode
def main():
if load_dotenv is not None:
load_dotenv()
env_chat_default = os.getenv("TC_CHAT_MODEL")
chat_default, offline_mode = parse_cli_args(sys.argv[1:])
if chat_default is None:
chat_default = env_chat_default
if not offline_mode:
offline_mode = os.getenv("TC_OFFLINE_MODE", "").strip().lower() in {"1", "true", "yes", "on"}
app = CheetahApp(
chat_default=chat_default,
offline_mode=offline_mode,
)
app.run()
if __name__ == "__main__":
main()