Skip to content

Commit e3888f3

Browse files
author
guozhen la
committed
init
1 parent a9f5d24 commit e3888f3

2 files changed

Lines changed: 150 additions & 24 deletions

File tree

flyteplugins/go/tasks/plugins/k8s/ray/ray.go

Lines changed: 29 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -79,8 +79,15 @@ func (rayJobResourceHandler) BuildResource(ctx context.Context, taskCtx pluginsC
7979
cfg := GetConfig()
8080
headReplicas := int32(1)
8181
headNodeRayStartParams := make(map[string]string)
82-
if rayJob.RayCluster.HeadGroupSpec != nil && rayJob.RayCluster.HeadGroupSpec.RayStartParams != nil {
83-
headNodeRayStartParams = rayJob.RayCluster.HeadGroupSpec.RayStartParams
82+
headGroupResources := &v1.ResourceRequirements{}
83+
if rayJob.RayCluster.HeadGroupSpec != nil{
84+
if rayJob.RayCluster.HeadGroupSpec.RayStartParams != nil {
85+
headNodeRayStartParams = rayJob.RayCluster.HeadGroupSpec.RayStartParams
86+
}
87+
headGroupResources, err = flytek8s.ToK8sResourceRequirements(rayJob.RayCluster.HeadGroupSpec.Resources)
88+
if err != nil {
89+
return nil, flyteerr.Errorf(flyteerr.BadTaskSpecification, "invalid TaskSpecification on Resources[%v], Err: [%v]", headGroupResources, err.Error())
90+
}
8491
} else if headNode := cfg.Defaults.HeadNode; len(headNode.StartParameters) > 0 {
8592
headNodeRayStartParams = headNode.StartParameters
8693
}
@@ -101,6 +108,10 @@ func (rayJobResourceHandler) BuildResource(ctx context.Context, taskCtx pluginsC
101108
headNodeRayStartParams[DisableUsageStatsStartParameter] = "true"
102109
}
103110

111+
if rayJob.RayCluster.Namespace != "" {
112+
objectMeta.Namespace = rayJob.RayCluster.Namespace
113+
}
114+
104115
enableIngress := true
105116
rayClusterSpec := rayv1alpha1.RayClusterSpec{
106117
HeadGroupSpec: rayv1alpha1.HeadGroupSpec{
@@ -114,7 +125,12 @@ func (rayJobResourceHandler) BuildResource(ctx context.Context, taskCtx pluginsC
114125
}
115126

116127
for _, spec := range rayJob.RayCluster.WorkerGroupSpec {
117-
workerPodTemplate := buildWorkerPodTemplate(&container, podSpec, objectMeta, taskCtx)
128+
workerGroupResources, err := flytek8s.ToK8sResourceRequirements(spec.Resources)
129+
if err != nil {
130+
return nil, flyteerr.Errorf(flyteerr.BadTaskSpecification, "invalid TaskSpecification on Resources[%v], Err: [%v]", workerGroupResources, err.Error())
131+
}
132+
133+
workerPodTemplate := buildWorkerPodTemplate(&container, podSpec, objectMeta, taskCtx, workerGroupResources)
118134

119135
minReplicas := spec.Replicas
120136
maxReplicas := spec.Replicas
@@ -153,7 +169,10 @@ func (rayJobResourceHandler) BuildResource(ctx context.Context, taskCtx pluginsC
153169
rayClusterSpec.WorkerGroupSpecs = append(rayClusterSpec.WorkerGroupSpecs, workerNodeSpec)
154170
}
155171

156-
serviceAccountName := flytek8s.GetServiceAccountNameFromTaskExecutionMetadata(taskCtx.TaskExecutionMetadata())
172+
serviceAccountName := rayJob.RayCluster.K8SServiceAccount
173+
if serviceAccountName == "" {
174+
serviceAccountName = flytek8s.GetServiceAccountNameFromTaskExecutionMetadata(taskCtx.TaskExecutionMetadata())
175+
}
157176

158177
rayClusterSpec.HeadGroupSpec.Template.Spec.ServiceAccountName = serviceAccountName
159178
for index := range rayClusterSpec.WorkerGroupSpecs {
@@ -180,12 +199,16 @@ func (rayJobResourceHandler) BuildResource(ctx context.Context, taskCtx pluginsC
180199
return &rayJobObject, nil
181200
}
182201

183-
func buildHeadPodTemplate(container *v1.Container, podSpec *v1.PodSpec, objectMeta *metav1.ObjectMeta, taskCtx pluginsCore.TaskExecutionContext) v1.PodTemplateSpec {
202+
func buildHeadPodTemplate(container *v1.Container, podSpec *v1.PodSpec, objectMeta *metav1.ObjectMeta, taskCtx pluginsCore.TaskExecutionContext, resources *v1.ResourceRequirements) v1.PodTemplateSpec {
184203
// Some configs are copy from https://github.com/ray-project/kuberay/blob/b72e6bdcd9b8c77a9dc6b5da8560910f3a0c3ffd/apiserver/pkg/util/cluster.go#L97
185204
// They should always be the same, so we could hard code here.
186205
primaryContainer := container.DeepCopy()
187206
primaryContainer.Name = "ray-head"
188207

208+
if len(resources.Requests) >= 1 || len(resources.Limits) >= 1 {
209+
primaryContainer.Resources = *resources
210+
}
211+
189212
envs := []v1.EnvVar{
190213
{
191214
Name: "MY_POD_IP",
@@ -232,7 +255,7 @@ func buildHeadPodTemplate(container *v1.Container, podSpec *v1.PodSpec, objectMe
232255
return podTemplateSpec
233256
}
234257

235-
func buildWorkerPodTemplate(container *v1.Container, podSpec *v1.PodSpec, objectMetadata *metav1.ObjectMeta, taskCtx pluginsCore.TaskExecutionContext) v1.PodTemplateSpec {
258+
func buildWorkerPodTemplate(container *v1.Container, podSpec *v1.PodSpec, objectMetadata *metav1.ObjectMeta, taskCtx pluginsCore.TaskExecutionContext, resources *v1.ResourceRequirements) v1.PodTemplateSpec {
236259
// Some configs are copy from https://github.com/ray-project/kuberay/blob/b72e6bdcd9b8c77a9dc6b5da8560910f3a0c3ffd/apiserver/pkg/util/cluster.go#L185
237260
// They should always be the same, so we could hard code here.
238261
initContainers := []v1.Container{

flyteplugins/go/tasks/plugins/k8s/ray/ray_test.go

Lines changed: 121 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -33,6 +33,8 @@ import (
3333

3434
const testImage = "image://"
3535
const serviceAccount = "ray_sa"
36+
const serviceAccountOverride = "ray_sa_override"
37+
const namespaceOverride = "ray_namespace_override"
3638

3739
var (
3840
dummyEnvVars = []*core.KeyValuePair{
@@ -43,6 +45,52 @@ var (
4345
"test-args",
4446
}
4547

48+
headResourceOverride = core.Resources{
49+
Requests: []*core.Resources_ResourceEntry{
50+
{
51+
Name: core.Resources_CPU,
52+
Value: "1000m",
53+
},
54+
{
55+
Name: core.Resources_MEMORY,
56+
Value: "2Gi",
57+
},
58+
},
59+
Limits: []*core.Resources_ResourceEntry{
60+
{
61+
Name: core.Resources_CPU,
62+
Value: "2000m",
63+
},
64+
{
65+
Name: core.Resources_MEMORY,
66+
Value: "4Gi",
67+
},
68+
},
69+
}
70+
71+
workerResourceOverride = core.Resources{
72+
Requests: []*core.Resources_ResourceEntry{
73+
{
74+
Name: core.Resources_CPU,
75+
Value: "5",
76+
},
77+
{
78+
Name: core.Resources_MEMORY,
79+
Value: "10G",
80+
},
81+
},
82+
Limits: []*core.Resources_ResourceEntry{
83+
{
84+
Name: core.Resources_CPU,
85+
Value: "10",
86+
},
87+
{
88+
Name: core.Resources_MEMORY,
89+
Value: "20G",
90+
},
91+
},
92+
}
93+
4694
resourceRequirements = &corev1.ResourceRequirements{
4795
Limits: corev1.ResourceList{
4896
corev1.ResourceCPU: resource.MustParse("1000m"),
@@ -68,6 +116,17 @@ func dummyRayCustomObj() *plugins.RayJob {
68116
}
69117
}
70118

119+
func dummyRayCustomObjWithOverrides() *plugins.RayJob {
120+
return &plugins.RayJob{
121+
RayCluster: &plugins.RayCluster{
122+
K8SServiceAccount: serviceAccountOverride,
123+
Namespace: namespaceOverride,
124+
HeadGroupSpec: &plugins.HeadGroupSpec{RayStartParams: map[string]string{"num-cpus": "1"}, Resources: &headResourceOverride},
125+
WorkerGroupSpec: []*plugins.WorkerGroupSpec{{GroupName: workerGroupName, Replicas: 3, Resources: &workerResourceOverride}},
126+
},
127+
}
128+
}
129+
71130
func dummyRayTaskTemplate(id string, rayJobObj *plugins.RayJob) *core.TaskTemplate {
72131

73132
ptObjJSON, err := utils.MarshalToString(rayJobObj)
@@ -172,26 +231,70 @@ func TestBuildResourceRay(t *testing.T) {
172231
assert.True(t, ok)
173232

174233
headReplica := int32(1)
175-
assert.Equal(t, ray.Spec.RayClusterSpec.HeadGroupSpec.Replicas, &headReplica)
176-
assert.Equal(t, ray.Spec.RayClusterSpec.HeadGroupSpec.Template.Spec.ServiceAccountName, serviceAccount)
177-
assert.Equal(t, ray.Spec.RayClusterSpec.HeadGroupSpec.RayStartParams,
178-
map[string]string{
179-
"dashboard-host": "0.0.0.0", "disable-usage-stats": "true", "include-dashboard": "true",
180-
"node-ip-address": "$MY_POD_IP", "num-cpus": "1"})
181-
assert.Equal(t, ray.Spec.RayClusterSpec.HeadGroupSpec.Template.Annotations, map[string]string{"annotation-1": "val1"})
182-
assert.Equal(t, ray.Spec.RayClusterSpec.HeadGroupSpec.Template.Labels, map[string]string{"label-1": "val1"})
183-
assert.Equal(t, ray.Spec.RayClusterSpec.HeadGroupSpec.Template.Spec.Tolerations, toleration)
234+
assert.Equal(t, &headReplica, ray.Spec.RayClusterSpec.HeadGroupSpec.Replicas)
235+
assert.Equal(t, serviceAccount, ray.Spec.RayClusterSpec.HeadGroupSpec.Template.Spec.ServiceAccountName)
236+
assert.Equal(t, map[string]string{"dashboard-host": "0.0.0.0", "include-dashboard": "true", "node-ip-address": "$MY_POD_IP", "num-cpus": "1"},
237+
ray.Spec.RayClusterSpec.HeadGroupSpec.RayStartParams)
238+
assert.Equal(t, map[string]string{"annotation-1": "val1"}, ray.Spec.RayClusterSpec.HeadGroupSpec.Template.Annotations)
239+
assert.Equal(t, map[string]string{"label-1": "val1"}, ray.Spec.RayClusterSpec.HeadGroupSpec.Template.Labels)
240+
assert.Equal(t, toleration, ray.Spec.RayClusterSpec.HeadGroupSpec.Template.Spec.Tolerations)
184241

185242
workerReplica := int32(3)
186-
assert.Equal(t, ray.Spec.RayClusterSpec.WorkerGroupSpecs[0].Replicas, &workerReplica)
187-
assert.Equal(t, ray.Spec.RayClusterSpec.WorkerGroupSpecs[0].MinReplicas, &workerReplica)
188-
assert.Equal(t, ray.Spec.RayClusterSpec.WorkerGroupSpecs[0].MaxReplicas, &workerReplica)
189-
assert.Equal(t, ray.Spec.RayClusterSpec.WorkerGroupSpecs[0].GroupName, workerGroupName)
190-
assert.Equal(t, ray.Spec.RayClusterSpec.WorkerGroupSpecs[0].Template.Spec.ServiceAccountName, serviceAccount)
191-
assert.Equal(t, ray.Spec.RayClusterSpec.WorkerGroupSpecs[0].RayStartParams, map[string]string{"disable-usage-stats": "true", "node-ip-address": "$MY_POD_IP"})
192-
assert.Equal(t, ray.Spec.RayClusterSpec.WorkerGroupSpecs[0].Template.Annotations, map[string]string{"annotation-1": "val1"})
193-
assert.Equal(t, ray.Spec.RayClusterSpec.WorkerGroupSpecs[0].Template.Labels, map[string]string{"label-1": "val1"})
194-
assert.Equal(t, ray.Spec.RayClusterSpec.WorkerGroupSpecs[0].Template.Spec.Tolerations, toleration)
243+
assert.Equal(t, &workerReplica, ray.Spec.RayClusterSpec.WorkerGroupSpecs[0].Replicas)
244+
assert.Equal(t, &workerReplica, ray.Spec.RayClusterSpec.WorkerGroupSpecs[0].MinReplicas)
245+
assert.Equal(t, &workerReplica, ray.Spec.RayClusterSpec.WorkerGroupSpecs[0].MaxReplicas)
246+
assert.Equal(t, workerGroupName, ray.Spec.RayClusterSpec.WorkerGroupSpecs[0].GroupName)
247+
assert.Equal(t, serviceAccount, ray.Spec.RayClusterSpec.WorkerGroupSpecs[0].Template.Spec.ServiceAccountName)
248+
assert.Equal(t, map[string]string{"node-ip-address": "$MY_POD_IP"}, ray.Spec.RayClusterSpec.WorkerGroupSpecs[0].RayStartParams)
249+
assert.Equal(t, map[string]string{"annotation-1": "val1"}, ray.Spec.RayClusterSpec.WorkerGroupSpecs[0].Template.Annotations)
250+
assert.Equal(t, map[string]string{"label-1": "val1"}, ray.Spec.RayClusterSpec.WorkerGroupSpecs[0].Template.Labels)
251+
assert.Equal(t, toleration, ray.Spec.RayClusterSpec.WorkerGroupSpecs[0].Template.Spec.Tolerations)
252+
}
253+
254+
func TestBuildResourceRayWithOverrides(t *testing.T) {
255+
rayJobResourceHandler := rayJobResourceHandler{}
256+
taskTemplate := dummyRayTaskTemplate("ray-id", dummyRayCustomObjWithOverrides())
257+
expectedHeadResources, _ := flytek8s.ToK8sResourceRequirements(&headResourceOverride)
258+
expectedWorkerResources, _ := flytek8s.ToK8sResourceRequirements(&workerResourceOverride)
259+
toleration := []corev1.Toleration{{
260+
Key: "storage",
261+
Value: "dedicated",
262+
Operator: corev1.TolerationOpExists,
263+
Effect: corev1.TaintEffectNoSchedule,
264+
}}
265+
err := config.SetK8sPluginConfig(&config.K8sPluginConfig{DefaultTolerations: toleration})
266+
assert.Nil(t, err)
267+
268+
RayResource, err := rayJobResourceHandler.BuildResource(context.TODO(), dummyRayTaskContext(taskTemplate))
269+
assert.Nil(t, err)
270+
271+
assert.NotNil(t, RayResource)
272+
ray, ok := RayResource.(*rayv1alpha1.RayJob)
273+
assert.True(t, ok)
274+
275+
headReplica := int32(1)
276+
assert.Equal(t, namespaceOverride, ray.Spec.RayClusterSpec.HeadGroupSpec.Template.ObjectMeta.Namespace)
277+
assert.Equal(t, &headReplica, ray.Spec.RayClusterSpec.HeadGroupSpec.Replicas)
278+
assert.Equal(t, serviceAccountOverride, ray.Spec.RayClusterSpec.HeadGroupSpec.Template.Spec.ServiceAccountName)
279+
assert.Equal(t, map[string]string{"dashboard-host": "0.0.0.0", "include-dashboard": "true", "node-ip-address": "$MY_POD_IP", "num-cpus": "1"},
280+
ray.Spec.RayClusterSpec.HeadGroupSpec.RayStartParams)
281+
assert.Equal(t, map[string]string{"annotation-1": "val1"}, ray.Spec.RayClusterSpec.HeadGroupSpec.Template.Annotations)
282+
assert.Equal(t, map[string]string{"label-1": "val1"}, ray.Spec.RayClusterSpec.HeadGroupSpec.Template.Labels)
283+
assert.Equal(t, toleration, ray.Spec.RayClusterSpec.HeadGroupSpec.Template.Spec.Tolerations)
284+
assert.Equal(t, *expectedHeadResources, ray.Spec.RayClusterSpec.HeadGroupSpec.Template.Spec.Containers[0].Resources)
285+
286+
workerReplica := int32(3)
287+
assert.Equal(t, namespaceOverride, ray.Spec.RayClusterSpec.WorkerGroupSpecs[0].Template.ObjectMeta.Namespace)
288+
assert.Equal(t, &workerReplica, ray.Spec.RayClusterSpec.WorkerGroupSpecs[0].Replicas)
289+
assert.Equal(t, &workerReplica, ray.Spec.RayClusterSpec.WorkerGroupSpecs[0].MinReplicas)
290+
assert.Equal(t, &workerReplica, ray.Spec.RayClusterSpec.WorkerGroupSpecs[0].MaxReplicas)
291+
assert.Equal(t, workerGroupName, ray.Spec.RayClusterSpec.WorkerGroupSpecs[0].GroupName)
292+
assert.Equal(t, serviceAccountOverride, ray.Spec.RayClusterSpec.WorkerGroupSpecs[0].Template.Spec.ServiceAccountName)
293+
assert.Equal(t, map[string]string{"node-ip-address": "$MY_POD_IP"}, ray.Spec.RayClusterSpec.WorkerGroupSpecs[0].RayStartParams)
294+
assert.Equal(t, map[string]string{"annotation-1": "val1"}, ray.Spec.RayClusterSpec.WorkerGroupSpecs[0].Template.Annotations)
295+
assert.Equal(t, map[string]string{"label-1": "val1"}, ray.Spec.RayClusterSpec.WorkerGroupSpecs[0].Template.Labels)
296+
assert.Equal(t, toleration, ray.Spec.RayClusterSpec.WorkerGroupSpecs[0].Template.Spec.Tolerations)
297+
assert.Equal(t, *expectedWorkerResources, ray.Spec.RayClusterSpec.WorkerGroupSpecs[0].Template.Spec.Containers[0].Resources)
195298
}
196299

197300
func TestDefaultStartParameters(t *testing.T) {

0 commit comments

Comments
 (0)