From 456f8218a2d167385ae5570b04d45fb1d73befa3 Mon Sep 17 00:00:00 2001 From: Jonas Knobloch Date: Fri, 10 Apr 2026 23:04:08 +0200 Subject: Construct dense with base --- tensor/dense.go | 16 ++++++++++++---- tensor/dense_test.go | 12 ++++++------ 2 files changed, 18 insertions(+), 10 deletions(-) (limited to 'tensor') diff --git a/tensor/dense.go b/tensor/dense.go index 723cba5..9a2deb6 100644 --- a/tensor/dense.go +++ b/tensor/dense.go @@ -11,7 +11,7 @@ type Dense[T any] struct { strides []int } -func NewDense[T any](shape []int) Dense[T] { +func NewDense[T any](shape []int, base []T) Dense[T] { r := len(shape) if r == 0 { @@ -33,6 +33,14 @@ func NewDense[T any](shape []int) Dense[T] { n *= shape[dim] } + if base == nil { + base = make([]T, n) + } + + if len(base) != n { + panic("unexpected base length") + } + strides := make([]int, len(shape)) for dim := range slices.Backward(strides) { @@ -46,7 +54,7 @@ func NewDense[T any](shape []int) Dense[T] { } out := Dense[T]{ - base: make([]T, n), + base: base, offset: 0, shape: make([]int, len(shape)), strides: strides, @@ -107,14 +115,14 @@ func (d Dense[T]) IsContiguous() bool { func (d Dense[T]) Contiguous() Dense[T] { if d.IsContiguous() { - out := NewDense[T](d.shape) + out := NewDense[T](d.shape, nil) copy(out.base, d.base[d.offset:d.offset+d.Size()]) return out } - out := NewDense[T](d.shape) + out := NewDense[T](d.shape, nil) idxs := make([]int, d.Rank()) diff --git a/tensor/dense_test.go b/tensor/dense_test.go index 5bc4e6f..93c5e00 100644 --- a/tensor/dense_test.go +++ b/tensor/dense_test.go @@ -6,7 +6,7 @@ import ( ) func TestNewDense_Strides(t *testing.T) { - d := NewDense[float32]([]int{2, 3, 4}) + d := NewDense[float32]([]int{2, 3, 4}, nil) expected := []int{12, 4, 1} @@ -16,7 +16,7 @@ func TestNewDense_Strides(t *testing.T) { } func TestDense_IsContiguous(t *testing.T) { - d := NewDense[float32]([]int{2, 3}) + d := NewDense[float32]([]int{2, 3}, nil) if !d.IsContiguous() { t.Fatalf("expected contiguous true") @@ -30,7 +30,7 @@ func TestDense_IsContiguous(t *testing.T) { } func TestDense_IsContiguousStrict(t *testing.T) { - d := NewDense[int]([]int{1, 3}) + d := NewDense[int]([]int{1, 3}, nil) v := Dense[int]{ base: d.base, @@ -45,7 +45,7 @@ func TestDense_IsContiguousStrict(t *testing.T) { } func TestDense_Contiguous(t *testing.T) { - d := NewDense[float32]([]int{2, 3}) + d := NewDense[float32]([]int{2, 3}, nil) permuted := d.Permute([]int{1, 0}) @@ -57,7 +57,7 @@ func TestDense_Contiguous(t *testing.T) { } func TestDense_Permute(t *testing.T) { - d := NewDense[float32]([]int{2, 3}) + d := NewDense[float32]([]int{2, 3}, nil) for i := range d.Size() { d.base[i] = float32(i) @@ -88,7 +88,7 @@ func TestDense_Permute(t *testing.T) { } func TestDense_All(t *testing.T) { - d := NewDense[float32]([]int{2, 3}) + d := NewDense[float32]([]int{2, 3}, nil) for i := range d.Size() { d.base[i] = float32(i) -- cgit v1.2.3