Compare commits
5 Commits
master
...
fix-serialize
| Author | SHA1 | Date | |
|---|---|---|---|
| 07dba26cd2 | |||
| e36a1785f6 | |||
| 5dde59a3ad | |||
| 57cbac78f4 | |||
| a9ed8bb401 |
@@ -61,10 +61,6 @@ jobs:
|
||||
os: [ubuntu-latest, windows-latest, macos-latest]
|
||||
python_version: ["3.12"]
|
||||
include:
|
||||
- os: windows-latest
|
||||
python_version: "3.7"
|
||||
- os: macos-latest
|
||||
python_version: "3.8"
|
||||
- os: ubuntu-latest
|
||||
python_version: "3.9"
|
||||
- os: windows-latest
|
||||
|
||||
+16
-4
@@ -30,9 +30,11 @@ from typing import (
|
||||
overload,
|
||||
)
|
||||
|
||||
import numpy
|
||||
import srsly
|
||||
from cymem.cymem import Pool
|
||||
from thinc.api import Config, CupyOps, Optimizer, get_current_ops
|
||||
from thinc.util import convert_recursive
|
||||
|
||||
from . import about, ty, util
|
||||
from .compat import Literal
|
||||
@@ -1212,7 +1214,7 @@ class Language:
|
||||
examples,
|
||||
):
|
||||
eg.predicted = doc
|
||||
return losses
|
||||
return _replace_numpy_floats(losses)
|
||||
|
||||
def rehearse(
|
||||
self,
|
||||
@@ -1463,7 +1465,7 @@ class Language:
|
||||
results = scorer.score(examples, per_component=per_component)
|
||||
n_words = sum(len(eg.predicted) for eg in examples)
|
||||
results["speed"] = n_words / (end_time - start_time)
|
||||
return results
|
||||
return _replace_numpy_floats(results)
|
||||
|
||||
def create_optimizer(self):
|
||||
"""Create an optimizer, usually using the [training.optimizer] config."""
|
||||
@@ -2141,7 +2143,9 @@ class Language:
|
||||
serializers["tokenizer"] = lambda p: self.tokenizer.to_disk( # type: ignore[union-attr]
|
||||
p, exclude=["vocab"]
|
||||
)
|
||||
serializers["meta.json"] = lambda p: srsly.write_json(p, self.meta)
|
||||
serializers["meta.json"] = lambda p: srsly.write_json(
|
||||
p, _replace_numpy_floats(self.meta)
|
||||
)
|
||||
serializers["config.cfg"] = lambda p: self.config.to_disk(p)
|
||||
for name, proc in self._components:
|
||||
if name in exclude:
|
||||
@@ -2255,7 +2259,9 @@ class Language:
|
||||
serializers: Dict[str, Callable[[], bytes]] = {}
|
||||
serializers["vocab"] = lambda: self.vocab.to_bytes(exclude=exclude)
|
||||
serializers["tokenizer"] = lambda: self.tokenizer.to_bytes(exclude=["vocab"]) # type: ignore[union-attr]
|
||||
serializers["meta.json"] = lambda: srsly.json_dumps(self.meta)
|
||||
serializers["meta.json"] = lambda: srsly.json_dumps(
|
||||
_replace_numpy_floats(self.meta)
|
||||
)
|
||||
serializers["config.cfg"] = lambda: self.config.to_bytes()
|
||||
for name, proc in self._components:
|
||||
if name in exclude:
|
||||
@@ -2306,6 +2312,12 @@ class Language:
|
||||
return self
|
||||
|
||||
|
||||
def _replace_numpy_floats(meta_dict: dict) -> dict:
|
||||
return convert_recursive(
|
||||
lambda v: isinstance(v, numpy.floating), lambda v: float(v), dict(meta_dict)
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class FactoryMeta:
|
||||
"""Dataclass containing information about a component and its defaults
|
||||
|
||||
Reference in New Issue
Block a user