summaryrefslogtreecommitdiff
path: root/tensor
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-04-10 23:04:08 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-04-10 23:04:08 +0200
commit456f8218a2d167385ae5570b04d45fb1d73befa3 (patch)
tree23b369ffc391ab0aca7ec842e90689fab098c7c0 /tensor
parentc1598d443f5bad23595de916c32553b620dcbeba (diff)
Construct dense with base
Diffstat (limited to 'tensor')
-rw-r--r--tensor/dense.go16
-rw-r--r--tensor/dense_test.go12
2 files changed, 18 insertions, 10 deletions
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)