Skip to content

Failing to train on example data: ValueError: Unbatching a tensor is only supported for rank >= 1 #3

Description

@robinmeyers

Hi there,

I'm trying to get rbpnet up and running on the example dataset. After creating a new conda env using python=3.8.10, I run pip install . from inside the rbpnet directory. Everything installs fine. However when I run the example:

rbpnet train --validation-data data/PTBP1_HepG2.data.tfrecord -c config.gin -d data/PTBP1_HepG2.dataspec.yml -o out/ data/PTBP1_HepG2.data.tfrecord

I get the following errror:

Traceback (most recent call last):
  File "/home/users/rmmeyers/miniconda3/envs/rbpnet/bin/rbpnet", line 8, in <module>
    sys.exit(main())
  File "/home/users/rmmeyers/miniconda3/envs/rbpnet/lib/python3.8/site-packages/click/core.py", line 1157, in __call__
    return self.main(*args, **kwargs)
  File "/home/users/rmmeyers/miniconda3/envs/rbpnet/lib/python3.8/site-packages/click/core.py", line 1078, in main
    rv = self.invoke(ctx)
  File "/home/users/rmmeyers/miniconda3/envs/rbpnet/lib/python3.8/site-packages/click/core.py", line 1688, in invoke
    return _process_result(sub_ctx.command.invoke(sub_ctx))
  File "/home/users/rmmeyers/miniconda3/envs/rbpnet/lib/python3.8/site-packages/click/core.py", line 1434, in invoke
    return ctx.invoke(self.callback, **ctx.params)
  File "/home/users/rmmeyers/miniconda3/envs/rbpnet/lib/python3.8/site-packages/click/core.py", line 783, in invoke
    return __callback(*args, **kwargs)
  File "/home/users/rmmeyers/miniconda3/envs/rbpnet/lib/python3.8/site-packages/rbpnet/bin/train.py", line 23, in main
    train(list(train_data), dataspec, config, output, val_data=validation_data)
  File "/home/users/rmmeyers/miniconda3/envs/rbpnet/lib/python3.8/site-packages/gin/config.py", line 1605, in gin_wrapper
    utils.augment_exception_message_and_reraise(e, err_str)
  File "/home/users/rmmeyers/miniconda3/envs/rbpnet/lib/python3.8/site-packages/gin/utils.py", line 41, in augment_exception_message_and_reraise
    raise proxy.with_traceback(exception.__traceback__) from None
  File "/home/users/rmmeyers/miniconda3/envs/rbpnet/lib/python3.8/site-packages/gin/config.py", line 1582, in gin_wrapper
    return fn(*new_args, **new_kwargs)
  File "/home/users/rmmeyers/miniconda3/envs/rbpnet/lib/python3.8/site-packages/rbpnet/train.py", line 116, in train
    val_dataset = val_data.dataset(batch_size=batch_size, shuffle=0, cache=cache)
  File "/home/users/rmmeyers/miniconda3/envs/rbpnet/lib/python3.8/site-packages/rbpnet/io.py", line 359, in dataset
    dataset = tf.data.Dataset.from_tensor_slices(self.tfrecords)
  File "/home/users/rmmeyers/miniconda3/envs/rbpnet/lib/python3.8/site-packages/tensorflow/python/data/ops/dataset_ops.py", line 831, in from_tensor_slices
    return from_tensor_slices_op._from_tensor_slices(tensors, name)
  File "/home/users/rmmeyers/miniconda3/envs/rbpnet/lib/python3.8/site-packages/tensorflow/python/data/ops/from_tensor_slices_op.py", line 25, in _from_tensor_slices
    return _TensorSliceDataset(tensors, name=name)
  File "/home/users/rmmeyers/miniconda3/envs/rbpnet/lib/python3.8/site-packages/tensorflow/python/data/ops/from_tensor_slices_op.py", line 38, in __init__
    self._structure = nest.map_structure(
  File "/home/users/rmmeyers/miniconda3/envs/rbpnet/lib/python3.8/site-packages/tensorflow/python/data/util/nest.py", line 122, in map_structure
    return nest_util.map_structure(
  File "/home/users/rmmeyers/miniconda3/envs/rbpnet/lib/python3.8/site-packages/tensorflow/python/util/nest_util.py", line 1056, in map_structure
    return _tf_data_map_structure(func, *structure, **kwargs)
  File "/home/users/rmmeyers/miniconda3/envs/rbpnet/lib/python3.8/site-packages/tensorflow/python/util/nest_util.py", line 1123, in _tf_data_map_structure
    return _tf_data_pack_sequence_as(structure[0], [func(*x) for x in entries])
  File "/home/users/rmmeyers/miniconda3/envs/rbpnet/lib/python3.8/site-packages/tensorflow/python/util/nest_util.py", line 1123, in <listcomp>
    return _tf_data_pack_sequence_as(structure[0], [func(*x) for x in entries])
  File "/home/users/rmmeyers/miniconda3/envs/rbpnet/lib/python3.8/site-packages/tensorflow/python/data/ops/from_tensor_slices_op.py", line 39, in <lambda>
    lambda component_spec: component_spec._unbatch(), batched_spec)  # pylint: disable=protected-access
  File "/home/users/rmmeyers/miniconda3/envs/rbpnet/lib/python3.8/site-packages/tensorflow/python/framework/tensor.py", line 429, in _unbatch
    raise ValueError("Unbatching a tensor is only supported for rank >= 1")
ValueError: Unbatching a tensor is only supported for rank >= 1
  In call to configurable 'train' (<function train at 0x7f9703889d30>)

Do you have any idea how to solve this issue? It looks like a problem reading in the tfrecord for the validation data.

Thanks,
Robin

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Projects

    No projects

    Milestone

    No milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions