ეს პოსტი გვიჩვენებს, თუ როგორ ხდება ინფერენსის გაშვება JAX / Flax-ის გამოყენებით. Stable Diffusion-ის მუშაობის დეტალებისთვის ან მისი GPU-ზე გაშვებისთვის, იხილეთ ეს Colab ნოუთბუქი. თუ გსურთ მიჰყვეთ ინსტრუქციებს, დააწკაპუნეთ ზემოთ მოცემულ ღილაკზე, რათა ეს პოსტი Colab ნოუთბუქის სახით გახსნათ. პირველ რიგში, დარწმუნდით, რომ იყენებთ TPU ბექენდს. თუ ამ ნოუთბუქს Colab-ში აწარმოებთ, ზემოთ მოცემულ მენიუში აირჩიეთ "Runtime", შემდეგ "Change runtime type" და შემდეგ "Hardware accelerator" პარამეტრებში აირჩიეთ TPU. გაითვალისწინეთ, რომ JAX არ არის ექსკლუზიური TPU-ებისთვის, მაგრამ ის საუკეთესოდ ამ აპარატურაზე მუშაობს, რადგან თითოეულ TPU სერვერს აქვს 8 TPU ამაჩქარებელი, რომლებიც პარალელურად მუშაობენ. Output: დარწმუნდით, რომ diffusers დაინსტალირებულია. შემდეგ შემოვიტანთ ყველა დამოკიდებულებას. მოდელის გამოყენებამდე, წონების ჩამოსატვირთად და გამოსაყენებლად საჭიროა მოდელის ლიცენზიის მიღება. ლიცენზია შექმნილია ასეთი ძლიერი მანქანური სწავლების სისტემის პოტენციური მავნე ზემოქმედების შესამცირებლად. მომხმარებლებს ვთხოვთ, სრულად და ყურადღებით გაეცნონ ლიცენზიას. გთავაზობთ შეჯამებას: Flax-ის წონები ხელმისაწვდომია Hugging Face Hub-ში, Stable Diffusion-ის რეპოზიტორიუმის ნაწილად. Stable Diffusion მოდელი ვრცელდება CreateML OpenRail-M ლიცენზიით. ეს არის ღია ლიცენზია, რომელიც არ აცხადებს უფლებებს თქვენ მიერ გენერირებულ გამოსახულებებზე და გიკრძალავთ განზრახ უკანონო ან მავნე შინაარსის წარმოებას. მოდელის ბარათი შეიცავს მეტ დეტალს, ამიტომ დაუთმეთ დრო მის წაკითხვას და ყურადღებით გაითვალისწინეთ, ეთანხმებით თუ არა ლიცენზიას. თუ ეთანხმებით, უნდა იყოთ რეგისტრირებული მომხმარებელი Hub-ში და გამოიყენოთ წვდომის ტოკენი, რომ კოდმა იმუშაოს. თქვენ გაქვთ წვდომის ტოკენის მიწოდების ორი ვარიანტი: შემდეგი უჯრა შესვლის ინტერფეისს წარმოადგენს, თუ ამ კომპიუტერზე უკვე არ გქონიათ ავთენტიფიკაცია. თქვენ დაგჭირდებათ თქვენი წვდომის ტოკენის ჩასმა. TPU მოწყობილობები მხარს უჭერენ bfloat16-ს, ეფექტურ ნახევრად-მცურავ ტიპს. ჩვენ მას გამოვიყენებთ ჩვენი ტესტებისთვის, მაგრამ შეგიძლიათ გამოიყენოთ float32 სრული სიზუსტისთვის. Flax არის ფუნქციური ფრეიმვორკი, ამიტომ მოდელები უსახელმწიფოებრივია (stateless) და პარამეტრები მათ გარეთ ინახება. წინასწარ გაწვრთნილი Flax-ის პაიპლაინის ჩატვირთვა დააბრუნებს როგორც თავად პაიპლაინს, ასევე მოდელის წონებს (ან პარამეტრებს). ჩვენ ვიყენებთ წონების bf16 ვერსიას, რაც იწვევს ტიპის გაფრთხილებებს, რომლებიც უსაფრთხოდ შეგიძლიათ უგულებელყოთ. ვინაიდან TPU-ებს ჩვეულებრივ 8 მოწყობილობა აქვთ პარალელურად მომუშავე, ჩვენს მოთხოვნას გავიმეორებთ იმდენჯერ, რამდენი მოწყობილობაც გვაქვს. შემდეგ ინფერენსს შევასრულებთ 8 მოწყობილობაზე ერთდროულად, თითოეული პასუხისმგებელია ერთი გამოსახულების გენერირებაზე. ამგვარად, ჩვენ მივიღებთ 8 გამოსახულებას იმავე დროში, რაც ერთ ჩიპს სჭირდება ერთი გამოსახულების გენერირებისთვის. მოთხოვნის გამეორების შემდეგ, ტოკენიზებული ტექსტის ID-ებს ვიღებთ პაიპლაინის `prepare_inputs` ფუნქციის გამოძახებით. ტოკენიზებული ტექსტის სიგრძე დაყენებულია 77 ტოკენზე, როგორც ამას მოითხოვს საბაზისო CLIP Text მოდელის კონფიგურაცია. Output: მოდელის პარამეტრები და შეყვანები უნდა განმეორდეს (replicated) 8 პარალელურ მოწყობილობაზე, რომლებიც გვაქვს. პარამეტრების ლექსიკონი რეპლიცირებულია `flax.jax_utils.replicate`-ის გამოყენებით, რომელიც გადის ლექსიკონში და ცვლის წონების ფორმას ისე, რომ ისინი 8-ჯერ განმეორდეს. მასივები რეპლიცირებულია `shard`-ის გამოყენებით. Output: ეს ფორმა ნიშნავს, რომ 8 მოწყობილობიდან თითოეული მიიღებს შეყვანად `jnp` მასივს `(1, 77)` ფორმით. ამდენად, 1 არის პაკეტის ზომა თითო მოწყობილობაზე. საკმარისი მეხსიერების მქონე TPU-ებში ის შეიძლება იყოს 1-ზე მეტი, თუ გვინდოდა ერთდროულად მრავალი გამოსახულების გენერირება (თითო ჩიპზე). ჩვენ თითქმის მზად ვართ გამოსახულებების გენერირებისთვის! უბრალოდ გვჭირდება შემთხვევითი რიცხვების გენერატორის შექმნა, რათა გადავცეთ გენერირების ფუნქციას. ეს არის სტანდარტული პროცედურა Flax-ში, რომელიც ძალიან სერიოზულია და კონკრეტული შეხედულებები აქვს შემთხვევით რიცხვებთან დაკავშირებით – ყველა ფუნქცია, რომელიც შემთხვევით რიცხვებს ეხება, ელის გენერატორის მიღებას. ეს უზრუნველყოფს რეპროდუცირებადობას, მაშინაც კი, როდესაც ვავარჯიშებთ მრავალ დისტრიბუციულ მოწყობილობაზე. ქვემოთ მოცემული დამხმარე ფუნქცია იყენებს "seed"-ს შემთხვევითი რიცხვების გენერატორის ინიციალიზაციისთვის. სანამ ერთსა და იმავე "seed"-ს ვიყენებთ, მივიღებთ ზუსტად იგივე შედეგებს. მოგვიანებით ნოუთბუქში შედეგების შესწავლისას თავისუფლად გამოიყენეთ სხვადასხვა "seed". ჩვენ ვიღებთ `rng`-ს და შემდეგ "ვყოფთ" მას 8-ჯერ, რათა თითოეულმა მოწყობილობამ მიიღოს განსხვავებული გენერატორი. შესაბამისად, თითოეული მოწყობილობა შექმნის განსხვავებულ გამოსახულებას, და მთლიანი პროცესი რეპროდუცირებადია. JAX კოდი შეიძლება კომპილირებული იყოს ეფექტურ წარმოდგენაში, რომელიც ძალიან სწრაფად მუშაობს. თუმცა, ჩვენ უნდა დავრწმუნდეთ, რომ ყველა შეყვანას აქვს იგივე ფორმა შემდგომ გამოძახებებში; წინააღმდეგ შემთხვევაში, JAX-ს მოუწევს კოდის ხელახლა კომპილირება, და ჩვენ ვერ ვისარგებლებდით ოპტიმიზებული სიჩქარით. Flax პაიპლაინს შეუძლია კოდის კომპილირება ჩვენთვის, თუ არგუმენტად გადავცემთ `jit = True`-ს. ის ასევე უზრუნველყოფს მოდელის პარალელურად გაშვებას 8 ხელმისაწვდომ მოწყობილობაზე. პირველად, როდესაც გავუშვებთ შემდეგ უჯრედს, კომპილაციას დიდი დრო დასჭირდება, მაგრამ შემდგომი გამოძახებები (თუნდაც სხვადასხვა შეყვანებით) ბევრად უფრო სწრაფი იქნება. მაგალითად, ჩემს ტესტირებისას TPU v2-8-ზე კომპილაციას ერთ წუთზე მეტი დასჭირდა, მაგრამ შემდეგში ინფერენსის გაშვებას დაახლოებით 7 წამი სჭირდება. Output: დაბრუნებულ მასივს აქვს ფორმა `(8, 1, 512, 512, 3)`. ჩვენ მას ვცვლით ფორმას მეორე განზომილების მოსაშორებლად და 8 გამოსახულების მისაღებად.