model downloads

This commit is contained in:
Sander Land committed 2020-06-21 20:06:15 +02:00
1 parent c4b13fabb4
commit 13698858ef
18 files changed
+274 -40

No files matched your search

+68 -28
View File
@@ -24,9 +24,10 @@ from katrain.core.constants import (
)
from katrain.core.engine import KataGoEngine
from katrain.core.lang import i18n
from katrain.core.utils import find_package_resource
from katrain.core.utils import find_package_resource, PATHS
from katrain.gui.kivyutils import BackgroundMixin, I18NSpinner
from katrain.gui.style import DEFAULT_FONT, EVAL_COLORS
from katrain.gui.widgets.progress_loader import ProgressLoader
class I18NPopup(Popup):
@@ -306,40 +307,79 @@ class AIPopup(QuickConfigGui):
for _ in range((self.max_options - len(mode_settings)) * 2):
self.options_grid.add_widget(Label())
class ConfigPopup(QuickConfigGui):
def __init__(self, katrain):
super().__init__(katrain)
self.paths = [self.katrain.config("engine/model"), "katrain/models", "~/.katrain"]
def build_and_set_properties(self, *_args):
super().build_and_set_properties()
# self.check_models()
def check_models(self, *args): # WIP
try:
model = self.collect_properties(self)["engine/model"]
except InputParseError:
self.model_files.values = []
return
done = set()
model_files = []
for path in self.paths + [self.model_path.text]:
path = path.rstrip("/\\")
if path.startswith("katrain"):
path = path.replace("katrain", PATHS["PACKAGE"].rstrip("/\\"), 1)
path = os.path.expanduser(path)
if not os.path.isdir(path):
path, _file = os.path.split(path)
slashpath = path.replace("\\", "/")
if slashpath in done or not os.path.isdir(path):
continue
done.add(slashpath)
files = [
f.replace("/", os.path.sep).replace(PATHS["PACKAGE"], "katrain")
for ftype in ["*.bin.gz", "*.txt.gz"]
for f in glob.glob(slashpath + "/" + ftype)
]
print(path,files)
if files and path not in self.paths:
self.paths.append(path) # persistent on paths with models found
model_files += files
models_available_msg = i18n._("models available").format(num=len(model_files))
self.model_files.values = [models_available_msg] + model_files
self.model_files.text = models_available_msg
if os.path.exists(model):
if os.path.isdir(model):
path = model.rstrip("\\/")
file = None
else:
if model.startswith("katrain"):
model = find_package_resource(model)
path, file = os.path.split(model)
files = sorted(
[
os.path.split(f)[1]
for ftype in ["*.bin.gz", "*.txt.gz"]
for f in glob.glob(path + os.path.sep + ftype)
]
)
self.model_files.values = files
print(file, files, file in files)
if file in files:
self.model_files.text = file
else:
self.model_files.values = []
MODELS = {
# "pure 20b": "https://github.com/lightvector/KataGo/releases/download/v1.4.0/g170-b20c256x2-s4384473088-d968438914.bin.gz",
# "pure 30b": "https://github.com/lightvector/KataGo/releases/download/v1.4.0/g170-b30c320x2-s3530176512-d968463914.bin.gz",
# "pure 40b": "https://github.com/lightvector/KataGo/releases/download/v1.4.0/g170-b40c256x2-s3708042240-d967973220.bin.gz",
"final 20b":"https://github.com/lightvector/KataGo/releases/download/v1.4.5/g170e-b20c256x2-s5303129600-d1228401921.bin.gz",
"final 30b": "https://github.com/lightvector/KataGo/releases/download/v1.4.5/g170-b30c320x2-s4824661760-d1229536699.bin.gz",
"final 40b":"https://github.com/lightvector/KataGo/releases/download/v1.4.5/g170-b40c256x2-s5095420928-d1229425124.bin.gz"
}
def download_models(self, *_largs):
def download_complete(req, tmp_path, path, model):
try:
os.rename(tmp_path, path)
self.katrain.log(f"Download of {model} model complete -> {path}", OUTPUT_INFO)
except Exception as e:
self.katrain.log(f"Download of {model} model complete, but could not move file: {e}", OUTPUT_ERROR)
self.check_models()
for name, url in self.MODELS.items():
filename = os.path.split(url)[1]
if not any(os.path.split(f)[1] == filename for f in self.model_files.values):
savepath = os.path.expanduser(os.path.join("~/.katrain", filename))
savepath_tmp = savepath + ".part"
self.katrain.log(f"Downloading {name} model from {url} to {savepath_tmp}", OUTPUT_INFO)
progress = ProgressLoader(
download_url=url,
path_to_file=savepath_tmp,
downloading_text=f"Downloading {name} model: " + "{}%",
download_complete=lambda req, tmp=savepath_tmp, path=savepath, model=name: download_complete(
req, tmp, path, model
),
download_redirected=lambda req: self.katrain.log(
f"Download {name} redirected {req.resp_headers}", OUTPUT_DEBUG
),
)
progress.start(self.download_progress_box)
def update_config(self, save_to_file=True):
updated = super().update_config(save_to_file=save_to_file)