Compare commits

...

5 Commits

Author SHA1 Message Date
Matthew Honnibal 07dba26cd2 Remove obsolete python versions from tests 2024-10-01 10:02:03 +02:00
Matthew Honnibal e36a1785f6 Add missing import 2024-10-01 09:57:46 +02:00
Matthew Honnibal 5dde59a3ad Format 2024-09-30 22:31:38 +02:00
Matthew Honnibal 57cbac78f4 Fix numpy floating values in meta.json for serialization 2024-09-30 22:26:08 +02:00
Matthew Honnibal a9ed8bb401 Replace numpy floats in evaluate and update 2024-09-30 22:22:41 +02:00
2 changed files with 16 additions and 8 deletions
-4
View File
@@ -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
View File
@@ -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