meta-pytorch / meta-pytorch/data

`insert_dp` for adding additional pipes (similar to `replace_dp` and `remove_dp`)

Open
#750 4 comments 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

Dominant language
Python
Stars
1.3k
Forks
179
Avg merge
6d 1h
Merged PRs (30d)
2

Description

🚀 The feature

Building on dataloader2.graph.replace_dp add a insert_dp possibly insert_dpg (insert a sub graph) functions to insert datapipes into an existing template.

Motivation, pitch

I want to be able to add a single or a graph of datapipes without replacing the existing datapipes.

My practical example is making a dqn multi-processing friendly async dqn

A default dqn agent has the following pipeline:

agent = AgentBase(model)
agent = StepFieldSelector(agent,field='state')
agent = SimpleModelRunner(agent,device=device)
agent = ArgMaxer(agent)
selector = EpsilonSelector(agent,min_epsilon=min_epsilon,max_epsilon=max_epsilon,max_steps=max_steps,device=device)
if logger_bases is not None: agent = EpsilonCollector(selector,logger_bases)
agent = ArgMaxer(agent,only_idx=True)
agent = NumpyConverter(agent)
agent = PyPrimativeConverter(agent)
agent = AgentHead(agent)

I want to make the tempalte / base dqn agent capable of syncing a model across spawn processes. So we insert a data pipe to sync the model.

agent = AgentBase(model)
agent = StepFieldSelector(agent,field='state')
#### agent = ModelSubscriber(agent,device=device) <- insert a pipe before the `SimpleModelRunner` pipe ####
agent = SimpleModelRunner(agent,device=device)
agent = ArgMaxer(agent)
selector = EpsilonSelector(agent,min_epsilon=min_epsilon,max_epsilon=max_epsilon,max_steps=max_steps,device=device)
if logger_bases is not None: agent = EpsilonCollector(selector,logger_bases)
agent = ArgMaxer(agent,only_idx=True)
agent = NumpyConverter(agent)
agent = PyPrimativeConverter(agent)
agent = AgentHead(agent)
Alternatives
Option 1

Add if statements / modify the template to contain most extensions, kind of like the EpsilonCollector I have above.

Option 2

Modify replace_db to support replacing a dp with a DataPipeGraph. So we would do something like:

agent_sub = ModelSubscriber(find_dps(agent,StepFieldSelector)[0],device=device) 
agent_sub = SimpleModelRunner(agent_sub,device=device)

replace_db(agent,SimpleModelRunner,agent_sub)
Additional context

Not super tested, but the implementation below I think can allow for inserting a DataPipe or an entire isolated DataPipeGraph

# I have a PassThroughIterPipe that acts as a location that the insert code can definitively know that
# it can reassign
class PassThroughIterPipe(dp.iter.IterDataPipe):
    def __init__(self,source_datapipe): self.source_datapipe = source_datapipe
    def __iter__(self): return (o for o in self.source_datapipe)

def find_dp(graph: DataPipeGraph, dp_type: Type[DataPipe]) -> DataPipe:
    pipes = find_dps(graph,dp_type)
    if len(pipes)==1: return pipes[0]
    elif len(pipes)>1:
        found_ids = set([id(pipe) for pipe in pipes])
        if len(found_ids)>1:
            warn(f"""There are {len(pipes)} pipes of type {dp_type}. If this is intended, 
                     please use `find_dps` directly. Returning first instance.""")
        return pipes[0]
    else:
        raise LookupError(f'Unable to find {dp_type} starting at {graph}')
    
find_dp.__doc__ = "Returns a single `DataPipe` as opposed to `find_dps`.\n"+find_dps.__doc__

def _insert_dp(recv_dp, send_graph: DataPipeGraph, old_dp: DataPipe, new_dp: DataPipe) -> None:
    old_dp_id = id(old_dp)
    for send_id in send_graph:
        if send_id == old_dp_id:
            # We do the same as replace_dp here by switching recv_dp to new_dp
            _assign_attr(recv_dp, old_dp, new_dp, inner_dp=True)
            
            # Replace the last datapipe in new_dp with the old_dp
            final_datapipe = find_dp(traverse(new_dp),PassThroughIterPipe)
            # But now we switch new_dp from the place holder pipe PassThroughIterPipe, to old_dp thus 
            # not breaking the chain. Havent tested if this works for whole graphs as new_dp
            _assign_attr(new_dp, final_datapipe, old_dp, inner_dp=True)
            # new_dp.source_datapipe
        else:
            send_dp, sub_send_graph = send_graph[send_id]
            _insert_dp(send_dp, sub_send_graph, old_dp, new_dp)

def insert_dp(graph: DataPipeGraph, on_datapipe: DataPipe, insert_datapipe: DataPipe) -> DataPipeGraph:
    r"""
    Given the graph of DataPipe generated by ``traverse`` function and the ``on_datapipe`` DataPipe to be reconnected and
    the new ``insert_datapipe`` DataPipe to be inserted after ``on_datapipe``, 
    return the new graph of DataPipe.
    """
    assert len(graph) == 1

    # Check if `on_datapipe` is that the head of the graph
    # If so, we `insert_datapipe`
    if id(on_datapipe) in graph: 
        graph = traverse(insert_datapipe, only_datapipe=True)

    final_datapipe = list(graph.values())[0][0]
    
    for recv_dp, send_graph in graph.values():
        _insert_dp(recv_dp, send_graph, on_datapipe, insert_datapipe)

    return traverse(final_datapipe, only_datapipe=True)

With the test being:

it_pipe = dp.iter.IterableWrapper([1,2,3,4,5,6])
pipe = it_pipe.cycle(count=2)
pipe = pipe.batch(batch_size=2)

new_dp = insert_dp(
    traverse(pipe,only_datapipe=True),
    find_dp(traverse(pipe,only_datapipe=True),dp.iter.Cycler),
    dp.iter.Header(PassThroughIterPipe([]),limit=4)
)

image

Contributor guide

Open the contributing guide

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

Start in torchdata/dataloader2/graph.py with replace_dp, remove_dp, traverse, and the proposed insert_dp entry point. Review the supplied PassThroughIterPipe and IterableWrapper/cycle/batch example, then determine how insertion should handle a single DataPipe versus an isolated graph. Done means the selected existing pipe remains connected and traversal produces the expected graph for the example.

Written by the indexing model from the issue text.

Assessment

Tech stack
python
Domain
data-engineering
Issue type
Feature
Difficulty
5/5
Estimated time
Over a week
Activity status
Stale
Clarity
Mostly clear
Newbie friendliness
35/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.