Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 2 additions & 16 deletions pkg/device/amd/device.go
Original file line number Diff line number Diff line change
Expand Up @@ -131,28 +131,14 @@ func (dev *AMDDevices) PatchAnnotations(pod *corev1.Pod, annoinput *map[string]s
}

func (dev *AMDDevices) LockNode(n *corev1.Node, p *corev1.Pod) error {
found := false
for _, val := range p.Spec.Containers {
if (dev.GenerateResourceRequests(&val).Nums) > 0 {
found = true
break
}
}
if !found {
if !device.PodRequiresDevice(dev, p) {
return nil
}
return nodelock.LockNode(n.Name, NodeLockAMD, p)
}

func (dev *AMDDevices) ReleaseNodeLock(n *corev1.Node, p *corev1.Pod) error {
found := false
for _, val := range p.Spec.Containers {
if (dev.GenerateResourceRequests(&val).Nums) > 0 {
found = true
break
}
}
if !found {
if !device.PodRequiresDevice(dev, p) {
return nil
}
return nodelock.ReleaseNodeLock(n.Name, NodeLockAMD, p, false)
Expand Down
18 changes: 2 additions & 16 deletions pkg/device/ascend/device.go
Original file line number Diff line number Diff line change
Expand Up @@ -256,29 +256,15 @@ func (dev *Devices) PatchAnnotations(pod *corev1.Pod, annoInput *map[string]stri
}

func (dev *Devices) LockNode(n *corev1.Node, p *corev1.Pod) error {
found := false
for _, val := range p.Spec.Containers {
if (dev.GenerateResourceRequests(&val).Nums) > 0 {
found = true
break
}
}
if !found {
if !device.PodRequiresDevice(dev, p) {
return nil
}

return nodelock.LockNode(n.Name, NodeLockAscend, p)
}

func (dev *Devices) ReleaseNodeLock(n *corev1.Node, p *corev1.Pod) error {
found := false
for _, val := range p.Spec.Containers {
if (dev.GenerateResourceRequests(&val).Nums) > 0 {
found = true
break
}
}
if !found {
if !device.PodRequiresDevice(dev, p) {
return nil
}

Expand Down
18 changes: 2 additions & 16 deletions pkg/device/biren/device.go
Original file line number Diff line number Diff line change
Expand Up @@ -96,28 +96,14 @@ func (dev *BirenDevices) MutateAdmission(ctr *corev1.Container, p *corev1.Pod) (
}

func (dev *BirenDevices) LockNode(n *corev1.Node, p *corev1.Pod) error {
found := false
for _, val := range p.Spec.Containers {
if (dev.GenerateResourceRequests(&val).Nums) > 0 {
found = true
break
}
}
if !found {
if !device.PodRequiresDevice(dev, p) {
return nil
}
return nodelock.LockNode(n.Name, nodelock.NodeLockKey, p)
}

func (dev *BirenDevices) ReleaseNodeLock(n *corev1.Node, p *corev1.Pod) error {
found := false
for _, val := range p.Spec.Containers {
if (dev.GenerateResourceRequests(&val).Nums) > 0 {
found = true
break
}
}
if !found {
if !device.PodRequiresDevice(dev, p) {
return nil
}
return nodelock.ReleaseNodeLock(n.Name, nodelock.NodeLockKey, p, false)
Expand Down
30 changes: 28 additions & 2 deletions pkg/device/biren/device_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -423,21 +423,47 @@ func TestDevices_LockNode(t *testing.T) {
expectError bool
}{
{
name: "Test with no containers",
name: "no containers — skip lock",
node: &corev1.Node{},
pod: &corev1.Pod{Spec: corev1.PodSpec{}},
hasLock: false,
expectError: false,
},
{
name: "Test with non-zero resource requests",
name: "regular-container GPU request — acquires lock",
node: &corev1.Node{},
pod: &corev1.Pod{Spec: corev1.PodSpec{Containers: []corev1.Container{{Resources: corev1.ResourceRequirements{Requests: corev1.ResourceList{
"birentech.com/gpu": resource.MustParse("1"),
}}}}}},
hasLock: true,
expectError: false,
},
{
name: "init-container-only GPU request — acquires lock",
node: &corev1.Node{},
pod: &corev1.Pod{Spec: corev1.PodSpec{
InitContainers: []corev1.Container{{Resources: corev1.ResourceRequirements{Requests: corev1.ResourceList{
"birentech.com/gpu": resource.MustParse("1"),
}}}},
Containers: []corev1.Container{{Name: "cpu-app"}},
}},
hasLock: true,
expectError: false,
},
{
name: "init+regular GPU request — acquires lock",
node: &corev1.Node{},
pod: &corev1.Pod{Spec: corev1.PodSpec{
InitContainers: []corev1.Container{{Resources: corev1.ResourceRequirements{Requests: corev1.ResourceList{
"birentech.com/gpu": resource.MustParse("1"),
}}}},
Containers: []corev1.Container{{Resources: corev1.ResourceRequirements{Requests: corev1.ResourceList{
"birentech.com/gpu": resource.MustParse("1"),
}}}},
}},
hasLock: true,
expectError: false,
},
}

for _, tt := range tests {
Expand Down
9 changes: 1 addition & 8 deletions pkg/device/cambricon/device.go
Original file line number Diff line number Diff line change
Expand Up @@ -123,14 +123,7 @@ func (dev *CambriconDevices) setNodeLock(node *corev1.Node) error {
}

func (dev *CambriconDevices) LockNode(n *corev1.Node, p *corev1.Pod) error {
found := false
for _, val := range p.Spec.Containers {
if (dev.GenerateResourceRequests(&val).Nums) > 0 {
found = true
break
}
}
if !found {
if !device.PodRequiresDevice(dev, p) {
return nil
}
if _, ok := n.Annotations[DsmluLockTime]; !ok {
Expand Down
19 changes: 19 additions & 0 deletions pkg/device/devices.go
Original file line number Diff line number Diff line change
Expand Up @@ -733,3 +733,22 @@ func CheckType(annos map[string]string, cardType, useKey, noUseKey string) bool
}
return true
}

// PodRequiresDevice returns true if any container (init container or regular container)
// in the pod requests resources from the specified device generator.
func PodRequiresDevice(dev Devices, p *corev1.Pod) bool {
if p == nil || dev == nil {
return false
}
for i := range p.Spec.InitContainers {
if dev.GenerateResourceRequests(&p.Spec.InitContainers[i]).Nums > 0 {
return true
}
}
for i := range p.Spec.Containers {
if dev.GenerateResourceRequests(&p.Spec.Containers[i]).Nums > 0 {
return true
}
}
return false
}
103 changes: 103 additions & 0 deletions pkg/device/devices_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -1869,3 +1869,106 @@ func TestDecodeNodeDevicesLegacyFormat(t *testing.T) {
},
}, decoded)
}

func TestPodRequiresDevice(t *testing.T) {
mockDev := &mockDevices{
resourceRequest: ContainerDeviceRequest{
Nums: 1,
Type: "NVIDIA",
Memreq: 1000,
Coresreq: 10,
},
}

gpuCtr := corev1.Container{
Name: "gpu-ctr",
Resources: corev1.ResourceRequirements{
Limits: corev1.ResourceList{
"nvidia.com/gpu": resource.MustParse("1"),
},
},
}
noGpuCtr := corev1.Container{
Name: "no-gpu-ctr",
Resources: corev1.ResourceRequirements{
Limits: corev1.ResourceList{
"cpu": resource.MustParse("1"),
},
},
}

tests := []struct {
name string
dev Devices
pod *corev1.Pod
want bool
}{
{
name: "nil pod returns false",
dev: mockDev,
pod: nil,
want: false,
},
{
name: "nil dev returns false",
dev: nil,
pod: &corev1.Pod{
Spec: corev1.PodSpec{
Containers: []corev1.Container{gpuCtr},
},
},
want: false,
},
{
name: "no device request in init or regular containers",
dev: mockDev,
pod: &corev1.Pod{
Spec: corev1.PodSpec{
InitContainers: []corev1.Container{noGpuCtr},
Containers: []corev1.Container{noGpuCtr},
},
},
want: false,
},
{
name: "regular-container-only request",
dev: mockDev,
pod: &corev1.Pod{
Spec: corev1.PodSpec{
InitContainers: []corev1.Container{noGpuCtr},
Containers: []corev1.Container{gpuCtr},
},
},
want: true,
},
{
name: "init-container-only request",
dev: mockDev,
pod: &corev1.Pod{
Spec: corev1.PodSpec{
InitContainers: []corev1.Container{gpuCtr},
Containers: []corev1.Container{noGpuCtr},
},
},
want: true,
},
{
name: "init and regular container requests",
dev: mockDev,
pod: &corev1.Pod{
Spec: corev1.PodSpec{
InitContainers: []corev1.Container{gpuCtr},
Containers: []corev1.Container{gpuCtr},
},
},
want: true,
},
}

for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := PodRequiresDevice(tt.dev, tt.pod)
assert.Equal(t, got, tt.want)
})
}
}
18 changes: 2 additions & 16 deletions pkg/device/hygon/device.go
Original file line number Diff line number Diff line change
Expand Up @@ -101,28 +101,14 @@ func checkDCUtype(annos map[string]string, cardtype string) bool {
}

func (dev *DCUDevices) LockNode(n *corev1.Node, p *corev1.Pod) error {
found := false
for _, val := range p.Spec.Containers {
if (dev.GenerateResourceRequests(&val).Nums) > 0 {
found = true
break
}
}
if !found {
if !device.PodRequiresDevice(dev, p) {
return nil
}
return nodelock.LockNode(n.Name, NodeLockDCU, p)
}

func (dev *DCUDevices) ReleaseNodeLock(n *corev1.Node, p *corev1.Pod) error {
found := false
for _, val := range p.Spec.Containers {
if (dev.GenerateResourceRequests(&val).Nums) > 0 {
found = true
break
}
}
if !found {
if !device.PodRequiresDevice(dev, p) {
return nil
}
return nodelock.ReleaseNodeLock(n.Name, NodeLockDCU, p, false)
Expand Down
Loading
Loading