mirror of
https://github.com/wassname/ray.git
synced 2026-08-13 12:30:18 +08:00
[rllib] Fix output API when lz4 not installed (#5421)
This commit is contained in:
@@ -19,7 +19,7 @@ from ray.rllib.policy.sample_batch import MultiAgentBatch
|
||||
from ray.rllib.offline.io_context import IOContext
|
||||
from ray.rllib.offline.output_writer import OutputWriter
|
||||
from ray.rllib.utils.annotations import override, PublicAPI
|
||||
from ray.rllib.utils.compression import pack
|
||||
from ray.rllib.utils.compression import pack, compression_supported
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -99,7 +99,7 @@ class JsonWriter(OutputWriter):
|
||||
|
||||
|
||||
def _to_jsonable(v, compress):
|
||||
if compress:
|
||||
if compress and compression_supported():
|
||||
return str(pack(v))
|
||||
elif isinstance(v, np.ndarray):
|
||||
return v.tolist()
|
||||
|
||||
@@ -23,6 +23,11 @@ except ImportError:
|
||||
LZ4_ENABLED = False
|
||||
|
||||
|
||||
@DeveloperAPI
|
||||
def compression_supported():
|
||||
return LZ4_ENABLED
|
||||
|
||||
|
||||
@DeveloperAPI
|
||||
def pack(data):
|
||||
if LZ4_ENABLED:
|
||||
|
||||
Reference in New Issue
Block a user