JingsongLi commented on code in PR #9708:
URL: https://github.com/apache/paimon/pull/9708#discussion_r3976716641
##########
paimon-python/pypaimon/multimodal/lerobot/writer.py:
##########
@@ -47,12 +63,204 @@
"index": {"dtype": "int64", "shape": (1,), "names": None},
"task_index": {"dtype": "int64", "shape": (1,), "names": None},
}
-_TASK_FEATURE = {"dtype": "string", "shape": (1,), "names": None}
_STATE_PREFIX = "pypaimon.lerobot."
_STATE_VERSION = _STATE_PREFIX + "state-version"
+_STATE_VERSION_VALUE = "1"
_NEXT_INDEX = _STATE_PREFIX + "next-index"
_NEXT_EPISODE_INDEX = _STATE_PREFIX + "next-episode-index"
-_TASK_INDICES = _STATE_PREFIX + "task-indices"
+_STAT_NAMES = (
+ "min", "max", "mean", "std", "count",
+ "q01", "q10", "q50", "q90", "q99",
+)
+_INTEGER_DTYPES = {
+ "int8", "int16", "int32", "int64", "uint8", "uint16", "uint32",
+}
+_EPISODE_STATE_COLUMNS = list(_EMPTY_EPISODES_SCHEMA.names) + [
+ "stats/index/count",
+]
+
+
+def _read_arrow(table, projection=None):
+ builder = table.new_read_builder()
+ if projection is not None:
+ builder = builder.with_projection(projection)
+ return builder.new_read().to_arrow(builder.new_scan().plan().splits())
+
+
+def _lerobot_stats_functions():
+ try:
+ from lerobot.datasets.compute_stats import (
+ aggregate_stats,
+ auto_downsample_height_width,
+ compute_episode_stats,
+ get_feature_stats,
+ sample_indices,
+ )
+ except ImportError as error:
+ raise ImportError(
+ "PaimonLeRobotWriter statistics require LeRobot; install "
+ "'pypaimon[lerobot]'.") from error
+ return (aggregate_stats, auto_downsample_height_width,
+ compute_episode_stats, get_feature_stats, sample_indices)
+
+
+def _nested_list_type(value_type, depth):
+ for _ in range(depth):
+ value_type = pa.list_(value_type)
+ return value_type
+
+
+def _episode_schema(features):
+ fields = list(_EMPTY_EPISODES_SCHEMA)
+ for name, feature in features.items():
+ dtype = str(feature.get("dtype", ""))
+ if dtype == "string":
+ continue
+ feature_shape = _feature_shape(feature, name)
+ for stat in _STAT_NAMES:
+ if stat == "count":
+ value_type = pa.int64()
+ depth = 1
+ elif dtype == "image":
+ value_type = pa.float64()
+ depth = 3
+ elif stat in ("min", "max") and dtype in _INTEGER_DTYPES:
+ value_type = pa.int64()
+ elif stat in ("min", "max") and dtype in ("bool", "boolean"):
+ value_type = pa.bool_()
+ else:
+ value_type = pa.float64()
+ if stat != "count" and dtype != "image":
+ depth = max(1, len(feature_shape))
+ fields.append(pa.field(
+ "stats/%s/%s" % (name, stat),
+ _nested_list_type(value_type, depth),
+ nullable=False,
+ ))
+ return pa.schema(fields)
+
+
+def _metadata_values(table, component):
+ result = {}
+ for row in _read_arrow(table).to_pylist():
+ if row["key"] in result:
+ raise ValueError(
+ "Existing LeRobot %s metadata repeats key %r."
+ % (component, row["key"]))
+ try:
+ result[row["key"]] = json.loads(row["value"])
+ except (TypeError, ValueError) as error:
+ raise ValueError(
+ "Existing LeRobot %s metadata is invalid."
+ % component) from error
+ return result
+
+
+def _image_stats(values):
+ try:
+ from PIL import Image
+ except ImportError as error:
+ raise ImportError(
+ "PaimonLeRobotWriter image statistics require Pillow from "
+ "'pypaimon[lerobot]'.") from error
+ (_, downsample, _, get_feature_stats,
+ sample_indices) = _lerobot_stats_functions()
+ images = []
+ for index in sample_indices(len(values)):
+ with Image.open(io.BytesIO(values[index])) as image:
+ array = np.asarray(image.convert("RGB"), dtype=np.uint8)
+ images.append(downsample(np.transpose(array, (2, 0, 1))))
+ stats = get_feature_stats(
+ np.stack(images), axis=(0, 2, 3), keepdims=True)
+ return {
+ name: value if name == "count" else np.squeeze(
+ value / 255.0, axis=0)
+ for name, value in stats.items()
+ }
+
+
+def _compute_stats(episode, features):
+ _, _, compute_episode_stats, _, _ = \
+ _lerobot_stats_functions()
+ data = {}
+ numeric_features = {}
+ reshaped_features = {}
+ result = {}
+ for name, feature in features.items():
+ dtype = str(feature.get("dtype", ""))
+ values = episode.column(name).to_pylist()
+ if dtype == "image":
+ result[name] = _image_stats(values)
+ else:
+ numeric_features[name] = feature
+ array = (
+ values if dtype == "string" else np.asarray(
+ values,
+ dtype=np.dtype("bool" if dtype == "boolean" else dtype),
+ )
+ )
+ if dtype != "string" and array.ndim > 2:
+ # Keep higher-rank stats stable across one- and multi-frame
+ # episodes while delegating the calculation to LeRobot.
+ reshaped_features[name] = array.shape[1:]
+ array = array.reshape(array.shape[0], -1)
+ data[name] = array
+ numeric_stats = compute_episode_stats(data, numeric_features)
+ for name, shape in reshaped_features.items():
+ numeric_stats[name] = {
+ stat: value if stat == "count" else value.reshape(shape)
+ for stat, value in numeric_stats[name].items()
+ }
+ result.update(numeric_stats)
+ return result
+
+
+def _aggregate_stats(stats_list, features):
+ aggregate_stats = _lerobot_stats_functions()[0]
+ result = {}
+ for name, feature in features.items():
+ if feature.get("dtype") == "string":
+ continue
+ key = "image" if feature.get("dtype") == "image" else "feature"
+ result[name] = aggregate_stats([
+ {key: stats[name]} for stats in stats_list
+ ])[key]
+ return result
+
+
+def _indexed_metadata_table(component, entries):
+ import pandas as pd
+
+ entries = list(entries)
+ indices, labels = zip(*entries) if entries else ((), ())
+ return pa.Table.from_pandas(pd.DataFrame(
+ {component + "_index": np.asarray(indices, dtype=np.int64)},
+ index=pd.Index(
+ labels, dtype="string", name=component,
Review Comment:
[P2] Pin the string storage used for task and subtask labels
`pd.Index(..., dtype="string")` follows the global
`pd.options.mode.string_storage` setting. With the supported pandas 2.3.3 /
PyArrow 19.0.1 combination and `mode.string_storage="pyarrow"`, both empty and
populated label indexes become Arrow `large_string`, which Paimon's schema
parser does not support. Creating a writer then fails with `Unsupported pyarrow
type: large_string` after creating part of the table group, and reopening an
existing table fails its companion-schema check. Changing this setting between
construction and `flush()` also causes a schema mismatch and makes the writer
terminal. I verified that the previous commit succeeds under the same setting.
Please pin the index dtype to `pd.StringDtype(storage="python")`, or otherwise
explicitly produce `pa.string()` while preserving pandas index metadata. The
explicit Python storage dtype passed create/write/resume and Unicode
task/subtask label round trips with the global setting still on `"pyarrow"`.
--
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.
To unsubscribe, e-mail: [email protected]
For queries about this service, please contact Infrastructure at:
[email protected]