Skip to content
Snippets Groups Projects
Select Git revision
  • 10df2868cfca9da09501592992b5efeedc004788
  • master default protected
  • v2.0-dev protected
  • zikeliml/Task-96-dotExporterForAST
  • zikeliml/124-rework-tutorials
  • fma
  • fhennig/v2.0-deprecations
  • holzer-master-patch-46757
  • 66-absolute-access-is-probably-not-copied-correctly-after-_eval_subs
  • gpu_bufferfield_fix
  • hyteg
  • vectorization_sqrt_fix
  • target_dh_refactoring
  • const_fix
  • improved_comm
  • gpu_liveness_opts
  • release/1.3.7 protected
  • release/1.3.6 protected
  • release/2.0.dev0 protected
  • release/1.3.5 protected
  • release/1.3.4 protected
  • release/1.3.3 protected
  • release/1.3.2 protected
  • release/1.3.1 protected
  • release/1.3 protected
  • release/1.2 protected
  • release/1.1.1 protected
  • release/1.1 protected
  • release/1.0.1 protected
  • release/1.0 protected
  • release/0.4.4 protected
  • last/Kerncraft
  • last/OpenCL
  • last/LLVM
  • release/0.4.3 protected
  • release/0.4.2 protected
36 results

__init__.py

Blame
  • torch_native_cpu.tmpl.cpp 2.12 KiB
    #include <torch/extension.h>
    
    #include <vector>
    
    using namespace pybind11::literals;
    
    using scalar_t = {{ dtype }};
    
    
    
    std::vector<at::Tensor> {{ kernel_name }}_forward(
    {%- for tensor in forward_tensors -%}
        at::Tensor {{ tensor }} {{- ", " if not loop.last -}}
    {%- endfor %})
    {
        //{% for tensor in forward_output_tensors -%}
        //auto {{tensor}} = at::zeros_like({{ forward_input_tensors[0] }});
        //{% endfor %}
    
        {% for i in dimensions -%}
        int _size_{{ forward_tensors[0] }}_{{ i }} = {{ forward_tensors[0] }}.size({{ i }});
        {% endfor %}
    
        {% for tensor in forward_tensors -%}
        {%- set last = loop.last -%}
        scalar_t* _data_{{ tensor }} = {{ tensor }}.data<scalar_t>();
        {% for i in dimensions -%}
        int _stride_{{tensor}}_{{i}} = {{tensor}}.strides()[{{ i }}];
        {% endfor -%}
        {% endfor -%}
    
        {{forward_kernel}}
    
        return {
        {%- for tensor in forward_output_tensors -%}
        {{ tensor }} {{- "," if not loop.last -}}
        {% endfor -%}
        };
    }
    
    std::vector<at::Tensor> {{ kernel_name }}_backward(
    {%- for tensor in backward_tensors -%}
        at::Tensor {{ tensor }} {{- ", " if not loop.last -}}
    {% endfor %})
    {
        //{% for tensor in backward_output_tensors -%}
        //auto {{tensor}} = at::zeros_like({{ backward_input_tensors[0] }});
        //{% endfor %}
    
        {% for tensor in backward_tensors -%}
        {%- set last = loop.last -%}
        scalar_t* _data_{{ tensor }} = {{ tensor }}.data<scalar_t>();
        {% for i in dimensions -%}
        int _stride_{{ tensor }}_{{i}} = {{ tensor }}.strides()[{{ i }}];
        {% endfor -%}
        {% endfor -%}
    
        {{backward_kernel}}
    
        return {
        {%- for tensor in backward_output_tensors -%}
        {{ tensor }} {{- "," if not loop.last -}}
        {% endfor -%}
        };
    }
    
    PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
      m.def("forward", &{{ kernel_name }}_forward, "{{ kernel_name }} forward (CPU)",
    {%- for tensor in forward_tensors -%}
        "{{ tensor }}"_a {{ ", " if not loop.last }}  
    {%- endfor -%} );
      m.def("backward", &{{ kernel_name }}_backward, "{{ kernel_name }} backward (CPU)",
    {%- for tensor in backward_tensors -%}
        "{{ tensor }}"_a {{ ", " if not loop.last }}  
    {%- endfor -%} );
    }