~/bend-docscommunity

tensor_band_concat.bend source

tensor_band_concat.bend on the hub · documented module

# Channel concatenation writes each input range directly into its final band.# Different input boundaries need no temporary dense tensor or aligned copy.import Baseimport ./tensor.bend as Tensorsimport ./tensor_band.bend as Bandsimport ./storage_buffer.bend as Storageimport ./traversal_partition.bend as Partitionimport ./storage_copy.bend as Copyimport ./tensor_band_schedule.bend as BandScheduledef geometry(shape: List<&2,U32>,axis: U32) -> Maybe<&2,Bands.Geometry>:  match shape axis:    case [channels,height,width] 0: Some{Bands.Geometry{1,channels,height,width}}    case [batch,channels,height,width] 1: Some{Bands.Geometry{batch,channels,height,width}}    case shape axis: None{}def tensor_planned(~Element: Data,geometry: Bands.Geometry,shape: List<&2,U32>,fill: Element,status: Tensors.Status,workers: U32,  owners: Bands.Bands(~Element) & Bands.Plan) -> Tensors.Tensor<Element> & Maybe<&2,Bands.Plan>:  (bands,layout) = owners  (Tensors.Banded{bands,geometry,shape,fill,status,workers},Some{layout})def tensor_plan(~Element: Data,tensor: Tensors.Tensor<Element>) -> Tensors.Tensor<Element> & Maybe<&2,Bands.Plan>:  match tensor:    case Tensors.Tensor{storage,shape,fill,status,workers}: (Tensors.Tensor{storage,shape,fill,status,workers},None{})    case Tensors.Banded{bands,geometry,shape,fill,status,workers}: tensor_planned(~Element,geometry,shape,fill,status,workers,Bands.plan(~Element,bands))def first_plan(first: Maybe<&2,Bands.Plan>,rest: Maybe<&2,Bands.Plan>) -> Maybe<&2,Bands.Plan>:  match first:    case Some{layout}: Some{layout}    case None{}: restdef single_plan(~Element: Data,first: Tensors.Tensor<Element> & Maybe<&2,Bands.Plan>) -> Tensors.Parts<Element> & Maybe<&2,Bands.Plan>:  (tensor,layout) = first  (Tensors.Single{tensor},layout)def cons_plan(~Element: Data,first: Tensors.Tensor<Element> & Maybe<&2,Bands.Plan>,rest: Tensors.Parts<Element> & Maybe<&2,Bands.Plan>) -> Tensors.Parts<Element> & Maybe<&2,Bands.Plan>:  (tensor,first) = first  (rest,other) = rest  (Tensors.Cons{tensor,rest},first_plan(first,other))def discover(~Element: Data,parts: Tensors.Parts<Element>) -> Tensors.Parts<Element> & Maybe<&2,Bands.Plan>:  match parts:    case Tensors.Single{tensor}: single_plan(~Element,tensor_plan(~Element,tensor))    case Tensors.Cons{tensor,rest}: cons_plan(~Element,tensor_plan(~Element,tensor),discover(~Element,rest))# Reuse the existing metadata traversal to classify the common band path.# Mixed representations and different plans use the dense concat caller.def compatible_plans(first: Maybe<&2,Bands.Plan>,rest: Maybe<&2,Bands.Plan>) -> Bool:  match first rest:    case Some{first} Some{rest}: Bands.same_plan(first,rest)    case first rest: False{}def single_compatible_plan(~Element: Data,tensor: Tensors.Tensor<Element>,layout: Maybe<&2,Bands.Plan>) ->  (Tensors.Parts<Element> & Maybe<&2,Bands.Plan>) & Bool:  match layout:    case Some{layout}: ((Tensors.Single{tensor},Some{layout}),True{})    case None{}: ((Tensors.Single{tensor},None{}),False{})def single_compatible(~Element: Data,first: Tensors.Tensor<Element> & Maybe<&2,Bands.Plan>) ->  (Tensors.Parts<Element> & Maybe<&2,Bands.Plan>) & Bool:  (tensor,layout) = first  single_compatible_plan(~Element,tensor,layout)def cons_compatible(~Element: Data,first: Tensors.Tensor<Element> & Maybe<&2,Bands.Plan>,  rest: (Tensors.Parts<Element> & Maybe<&2,Bands.Plan>) & Bool) -> (Tensors.Parts<Element> & Maybe<&2,Bands.Plan>) & Bool:  (tensor,first_layout) = first  (rest_and_layout,compatible) = rest  (rest_parts,rest_layout) = rest_and_layout  +first_layout = {first_layout : Maybe<&2,Bands.Plan>}  +rest_layout = {rest_layout : Maybe<&2,Bands.Plan>}  ((Tensors.Cons{tensor,rest_parts},first_plan(first_layout,rest_layout)),compatible && compatible_plans(first_layout,rest_layout))def discover_compatible(~Element: Data,parts: Tensors.Parts<Element>) -> (Tensors.Parts<Element> & Maybe<&2,Bands.Plan>) & Bool:  match parts:    case Tensors.Single{tensor}: single_compatible(~Element,tensor_plan(~Element,tensor))    case Tensors.Cons{tensor,rest}: cons_compatible(~Element,tensor_plan(~Element,tensor),discover_compatible(~Element,rest))def allocate(~Element: Data,layout: Bands.Plan,+geometry: Bands.Geometry,+fill: Element) -> Bands.Bands(~Element):  match layout:    case Bands.Rows{first,+rows}: Partition.Leaf{Tensors.buffer(~Element,fill,Bands.local_count(geometry,rows)),first,rows,Unit{}}    case Bands.Split{left,right}: Partition.Fork{allocate(~Element,left,geometry,fill),allocate(~Element,right,geometry,fill)}def dense_rows(~Element: Data,geometry: Bands.Geometry,+first: U32,+rows: U32,channel: U32,input: Storage.Buffer<Element>,output: Storage.Buffer<Element>) -> Storage.Buffer<Element> & Storage.Buffer<Element>:  Bands.Geometry{batch,+channels,+height,+width} = geometry  Bands.planes(~Element,U32.to_nat(channels),0,(rows * width : U32),(height * width : U32),(rows * width : U32),    ((first / height * channels * height + first % height) * width : U32),(channel * rows * width : U32),(input,output))def channel_offset(geometry: Bands.Geometry,channel: U32,rows: U32) -> U32:  Bands.Geometry{batch,channels,height,width} = geometry  (channel * rows * width : U32)def rows(~Element: Data,geometry: Bands.Geometry,+first: U32,+rows: U32,channel: U32,tensor: Tensors.Tensor<Element>,output: Storage.Buffer<Element>) -> Tensors.Tensor<Element> & Storage.Buffer<Element>:  match tensor:    case Tensors.Tensor{input,shape,fill,status,workers}:      Storage.bind_pair(~Storage.Buffer<Element>,~Storage.Buffer<Element>,~(Tensors.Tensor<Element> & Storage.Buffer<Element>),        dense_rows(~Element,geometry,first,rows,channel,input,output),input => output => (Tensors.Tensor{input,shape,fill,status,workers},output))    case Tensors.Banded{bands,+geometry,shape,fill,status,workers}:      Storage.bind_pair(~Bands.Bands(~Element),~Storage.Buffer<Element>,~(Tensors.Tensor<Element> & Storage.Buffer<Element>),        Bands.range_into(~Element,geometry,first,rows,channel_offset(geometry,channel,rows),bands,output),bands => output =>          (Tensors.Banded{bands,geometry,shape,fill,status,workers},output))def append_tree(~Element: Data,output: Bands.Bands(~Element),+geometry: Bands.Geometry,+channel: U32,input: Tensors.Tensor<Element>) -> Tensors.Tensor<Element> & Bands.Bands(~Element):  match output:    case Partition.Leaf{output,+first,+count,details}:      Storage.bind_pair(~Tensors.Tensor<Element>,~Storage.Buffer<Element>,~(Tensors.Tensor<Element> & Bands.Bands(~Element)),        rows(~Element,geometry,first,count,channel,input,output),input => output => (input,Partition.Leaf{output,first,count,Unit{}}))    case Partition.Fork{left,right}:      Storage.bind_pair(~Tensors.Tensor<Element>,~Bands.Bands(~Element),~(Tensors.Tensor<Element> & Bands.Bands(~Element)),        append_tree(~Element,left,geometry,channel,input),input => left =>          Storage.bind_pair(~Tensors.Tensor<Element>,~Bands.Bands(~Element),~(Tensors.Tensor<Element> & Bands.Bands(~Element)),            append_tree(~Element,right,geometry,channel,input),input => right => (input,Partition.Fork{left,right})))def channels(shape: List<&2,U32>) -> U32:  match shape:    case [channels,height,width]: channels    case [batch,channels,height,width]: channels    case shape: 0type Channel is Data:  Channel{geometry: Bands.Geometry,channel: U32}def append_band(~Element: Data,context: Channel,first: U32,rows: U32,output: Storage.Buffer<Element>,input: Storage.Buffer<Element>,resource: Unit) -> Bands.Bands(~Element):  Channel{+geometry,channel} = context  +rows = rows  Partition.Leaf{Bands.destination(~Element,Copy.range(~Element,Bands.local_count(geometry,rows),0,channel_offset(geometry,channel,rows),input,output)),first,rows,Unit{}}def append_planned(~Element: Data,+parallel: Bool,+geometry: Bands.Geometry,channel: U32,input: Bands.Bands(~Element),planned: Bands.Bands(~Element) & Bands.Plan) -> Bands.Bands(~Element):  (output,+layout) = planned  Bands.each(~Bands.Paired<Storage.Buffer<Element>,Storage.Buffer<Element>>,~Storage.Buffer<Element>,~Unit,~Channel,~Bands.unshared,    ~Bands.paired(~Storage.Buffer<Element>,~Storage.Buffer<Element>,~Unit,~Channel,~append_band(~Element)),    ~Bands.left_of(~Storage.Buffer<Element>,~Storage.Buffer<Element>),    layout,parallel && Bands.same_image(geometry,layout),parallel,geometry,Channel{geometry,channel},Bands.zip(~Storage.Buffer<Element>,~Storage.Buffer<Element>,output,input),Unit{})# Every output leaf receives the same rows of the input, through the shared# band recursion on the plan of the output. Trees of another shape keep the output.def append_aligned(~Element: Data,parallel: Bool,geometry: Bands.Geometry,channel: U32,input: Bands.Bands(~Element),output: Bands.Bands(~Element)) -> Bands.Bands(~Element):  append_planned(~Element,parallel,geometry,channel,input,Bands.plan(~Element,output))def appended_rows(~Element: Data,owners: Tensors.Tensor<Element> & Bands.Bands(~Element)) -> Bands.Bands(~Element):  (source,output) = owners  outputdef selected_rows(~Element: Data,same: Bool,+geometry: Bands.Geometry,channel: U32,workers: U32,layout: Bands.Plan,input: Bands.Bands(~Element),output: Bands.Bands(~Element),shape: List<&2,U32>,fill: Element,status: Tensors.Status,previous: U32) -> Bands.Bands(~Element):  match same:    case True{}: append_aligned(~Element,BandSchedule.parallel(geometry,workers,8.0,0.0,layout),geometry,channel,input,output)    case False{}:      appended_rows(~Element,append_tree(~Element,output,geometry,channel,Tensors.Banded{input,geometry,shape,fill,status,previous}))def aligned_or_rows(~Element: Data,+geometry: Bands.Geometry,channel: U32,workers: U32,  input: Bands.Bands(~Element) & Bands.Plan,output: Bands.Bands(~Element) & Bands.Plan,shape: List<&2,U32>,fill: Element,status: Tensors.Status,previous: U32) -> Bands.Bands(~Element):  (input,source_plan) = input  (output,+target_plan) = output  selected_rows(~Element,Bands.same_plan(source_plan,target_plan),geometry,channel,workers,target_plan,input,output,shape,fill,status,previous)def append_storage(~Element: Data,+geometry: Bands.Geometry,channel: U32,workers: U32,input: Tensors.Tensor<Element>,output: Bands.Bands(~Element)) -> Bands.Bands(~Element):  match input:    case Tensors.Tensor{storage,shape,fill,status,previous}:      appended_rows(~Element,append_tree(~Element,output,geometry,channel,Tensors.Tensor{storage,shape,fill,status,previous}))    case Tensors.Banded{bands,layout,shape,fill,status,previous}:      aligned_or_rows(~Element,geometry,channel,workers,Bands.plan(~Element,bands),Bands.plan(~Element,output),shape,fill,status,previous)def append_input(~Element: Data,geometry: Bands.Geometry,+channel: U32,+count: U32,workers: U32,input: Tensors.Tensor<Element>,output: Bands.Bands(~Element)) -> Bands.Bands(~Element) & U32:  Bands.Geometry{batch,channels,height,width} = geometry  (append_storage(~Element,Bands.Geometry{batch,count,height,width},channel,workers,input,output),(channel + count : U32))def append_part(~Element: Data,geometry: Bands.Geometry,workers: U32,input: Tensors.Tensor<Element>,state: Bands.Bands(~Element) & U32) -> Bands.Bands(~Element) & U32:  (output,channel) = state  Tensors.with_metadata(~Element,~(Bands.Bands(~Element) & U32),input,input => shape => fill => status => previous =>    append_input(~Element,geometry,channel,channels(shape),workers,input,output))def output(~Element: Data,state: Bands.Bands(~Element) & U32) -> Bands.Bands(~Element):  (output,channel) = state  outputdef append_parts(~Element: Data,parts: Tensors.Parts<Element>,+geometry: Bands.Geometry,+workers: U32,state: Bands.Bands(~Element) & U32) -> Bands.Bands(~Element):  match parts:    case Tensors.Single{input}: output(~Element,append_part(~Element,geometry,workers,input,state))    case Tensors.Cons{input,rest}: append_parts(~Element,rest,geometry,workers,append_part(~Element,geometry,workers,input,state))def join(~Element: Data,layout: Bands.Plan,+geometry: Bands.Geometry,parts: Tensors.Parts<Element>,shape: List<&2,U32>,+fill: Element,+workers: U32) -> Tensors.Tensor<Element>:  Tensors.Banded{append_parts(~Element,parts,geometry,workers,(allocate(~Element,layout,geometry,fill),0)),geometry,shape,fill,Tensors.Usable{},workers}